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 sqlmodel import Session
|
||||||
|
|
||||||
from agenteval.model_gateway import ModelGateway
|
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 (
|
from agenteval.services.model_configs import (
|
||||||
ModelConfigError,
|
ModelConfigError,
|
||||||
ModelConfigService,
|
ModelConfigService,
|
||||||
@ -52,6 +52,20 @@ def resolve_analysis_model(campaign: Campaign, session: Session) -> Optional[Mod
|
|||||||
return None
|
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(
|
def collect_failure_samples(
|
||||||
campaign_id: str,
|
campaign_id: str,
|
||||||
session: Session,
|
session: Session,
|
||||||
@ -76,12 +90,14 @@ def collect_failure_samples(
|
|||||||
turn = turns.get(result.turn_id)
|
turn = turns.get(result.turn_id)
|
||||||
user = extract_reply_text(turn.get_sent_message().get("msgBody")) if turn else ""
|
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 ""
|
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 "",
|
"run_id": run.id or "",
|
||||||
"user": user[:text_limit],
|
"user": user[:text_limit],
|
||||||
"reply": reply[:text_limit],
|
"reply": reply[:text_limit],
|
||||||
"reason": (result.reason or "")[:text_limit],
|
"reason": (result.reason or "")[:text_limit],
|
||||||
})
|
}
|
||||||
|
)
|
||||||
return {sid: items for sid, items in samples.items() if items}
|
return {sid: items for sid, items in samples.items() if items}
|
||||||
|
|
||||||
|
|
||||||
@ -123,10 +139,12 @@ async def _analyze_scenario(
|
|||||||
ensure_ascii=False,
|
ensure_ascii=False,
|
||||||
)
|
)
|
||||||
parsed = _parse_stage(
|
parsed = _parse_stage(
|
||||||
await chat_client([
|
await chat_client(
|
||||||
|
[
|
||||||
{"role": "system", "content": system_prompt},
|
{"role": "system", "content": system_prompt},
|
||||||
{"role": "user", "content": user_prompt},
|
{"role": "user", "content": user_prompt},
|
||||||
]),
|
]
|
||||||
|
),
|
||||||
f"场景「{entry.get('scenario_name', entry['scenario_id'])}」阶段一",
|
f"场景「{entry.get('scenario_name', entry['scenario_id'])}」阶段一",
|
||||||
)
|
)
|
||||||
narrative = parsed.get("narrative")
|
narrative = parsed.get("narrative")
|
||||||
@ -169,10 +187,12 @@ async def _synthesize(
|
|||||||
payload["探索发现"] = exploration_summary
|
payload["探索发现"] = exploration_summary
|
||||||
user_prompt = json.dumps(payload, ensure_ascii=False)
|
user_prompt = json.dumps(payload, ensure_ascii=False)
|
||||||
parsed = _parse_stage(
|
parsed = _parse_stage(
|
||||||
await chat_client([
|
await chat_client(
|
||||||
|
[
|
||||||
{"role": "system", "content": system_prompt},
|
{"role": "system", "content": system_prompt},
|
||||||
{"role": "user", "content": user_prompt},
|
{"role": "user", "content": user_prompt},
|
||||||
]),
|
]
|
||||||
|
),
|
||||||
"阶段二综合研判",
|
"阶段二综合研判",
|
||||||
)
|
)
|
||||||
overall = parsed.get("overall")
|
overall = parsed.get("overall")
|
||||||
@ -200,10 +220,9 @@ async def analyze_campaign(
|
|||||||
if not capability:
|
if not capability:
|
||||||
raise AnalysisError("活动没有可分析的场景数据")
|
raise AnalysisError("活动没有可分析的场景数据")
|
||||||
|
|
||||||
stage1 = await asyncio.gather(*[
|
stage1 = await asyncio.gather(
|
||||||
_analyze_scenario(entry, failure_samples.get(entry["scenario_id"], []), chat_client)
|
*[_analyze_scenario(entry, failure_samples.get(entry["scenario_id"], []), chat_client) for entry in capability]
|
||||||
for entry in capability
|
)
|
||||||
])
|
|
||||||
stage2 = await _synthesize(campaign, report, list(stage1), chat_client, exploration_summary=exploration_summary)
|
stage2 = await _synthesize(campaign, report, list(stage1), chat_client, exploration_summary=exploration_summary)
|
||||||
|
|
||||||
valid_scenario_ids = {entry["scenario_id"] for entry in capability}
|
valid_scenario_ids = {entry["scenario_id"] for entry in capability}
|
||||||
@ -212,13 +231,15 @@ async def analyze_campaign(
|
|||||||
if not isinstance(p, dict):
|
if not isinstance(p, dict):
|
||||||
continue
|
continue
|
||||||
severity = p.get("severity")
|
severity = p.get("severity")
|
||||||
problems.append({
|
problems.append(
|
||||||
|
{
|
||||||
"severity": severity if severity in _VALID_SEVERITIES else "medium",
|
"severity": severity if severity in _VALID_SEVERITIES else "medium",
|
||||||
"title": str(p.get("title", "")),
|
"title": str(p.get("title", "")),
|
||||||
"description": str(p.get("description", "")),
|
"description": str(p.get("description", "")),
|
||||||
"scenario_ids": [s for s in p.get("scenario_ids") or [] if s in valid_scenario_ids],
|
"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],
|
"evidence_run_ids": [r for r in p.get("evidence_run_ids") or [] if r in valid_run_ids],
|
||||||
})
|
}
|
||||||
|
)
|
||||||
suggestions = [
|
suggestions = [
|
||||||
{"priority": int(s.get("priority", i + 1)), "text": str(s.get("text", ""))}
|
{"priority": int(s.get("priority", i + 1)), "text": str(s.get("text", ""))}
|
||||||
for i, s in enumerate(stage2.get("suggestions") or [])
|
for i, s in enumerate(stage2.get("suggestions") or [])
|
||||||
@ -227,9 +248,7 @@ async def analyze_campaign(
|
|||||||
return {
|
return {
|
||||||
"overall": stage2["overall"],
|
"overall": stage2["overall"],
|
||||||
"problems": problems,
|
"problems": problems,
|
||||||
"scenario_narratives": [
|
"scenario_narratives": [{"scenario_id": s["scenario_id"], "narrative": s["narrative"]} for s in stage1],
|
||||||
{"scenario_id": s["scenario_id"], "narrative": s["narrative"]} for s in stage1
|
|
||||||
],
|
|
||||||
"suggestions": suggestions,
|
"suggestions": suggestions,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@ -5,16 +5,19 @@ from typing import Any, Optional
|
|||||||
from sqlmodel import Session
|
from sqlmodel import Session
|
||||||
|
|
||||||
from agenteval.evaluation.campaign_scheduler import clock_offset, elapsed_seconds
|
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 (
|
from agenteval.evaluation.report import (
|
||||||
build_campaign_timeline,
|
build_campaign_timeline,
|
||||||
generate_campaign_report,
|
generate_campaign_report,
|
||||||
|
load_campaign_report,
|
||||||
summarize_campaign_progress,
|
summarize_campaign_progress,
|
||||||
)
|
)
|
||||||
from agenteval.exploration.summary import summarize_campaign_exploration
|
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.db import iso_utc, utc_now
|
||||||
|
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignAnalysisRepository,
|
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
RunRepository,
|
RunRepository,
|
||||||
ScenarioRepository,
|
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:
|
class CampaignReadModel:
|
||||||
"""One interface for Campaign list, detail, report and export projections."""
|
"""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 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 complete_campaign
|
||||||
from agenteval.evaluation.campaign_lifecycle import start_campaign as start_campaign_lifecycle
|
from agenteval.evaluation.campaign_lifecycle import start_campaign as start_campaign_lifecycle
|
||||||
from agenteval.evaluation.campaign_scheduler import (
|
from agenteval.evaluation.campaign_scheduler import (
|
||||||
@ -33,7 +32,7 @@ from agenteval.evaluation.campaign_scheduler import (
|
|||||||
resolve_finalize,
|
resolve_finalize,
|
||||||
)
|
)
|
||||||
from agenteval.evaluation.engine import EvalEngine
|
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 (
|
from agenteval.models import (
|
||||||
Campaign,
|
Campaign,
|
||||||
CampaignStatus,
|
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
|
failure all skip silently — the analysis is an enhancement and must never
|
||||||
block or break campaign completion.
|
block or break campaign completion.
|
||||||
"""
|
"""
|
||||||
if campaign.time_scale != 1 or not campaign.id:
|
if not auto_intelligence_eligible(campaign, session):
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
if resolve_analysis_model(campaign, session) is None:
|
|
||||||
return
|
|
||||||
enqueue_campaign_analysis(campaign.id, triggered_by="auto")
|
enqueue_campaign_analysis(campaign.id, triggered_by="auto")
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
_logger.warning("活动 %s 自动分析触发失败(已跳过): %s", campaign.id, exc)
|
_logger.warning("活动 %s 自动分析触发失败(已跳过): %s", campaign.id, exc)
|
||||||
|
|||||||
@ -13,15 +13,10 @@ from typing import Any, Optional
|
|||||||
from sqlmodel import Session
|
from sqlmodel import Session
|
||||||
|
|
||||||
from agenteval.evaluation.analysis import ChatClient, resolve_analysis_model
|
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.models import Campaign, CampaignStatus
|
||||||
from agenteval.storage.db import iso_utc, utc_now
|
from agenteval.storage.async_job_repository import CampaignAnalysisRepository
|
||||||
from agenteval.storage.model_config_repository import ModelConfigRepository
|
from agenteval.storage.db import utc_now
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import CampaignRepository
|
||||||
CampaignAnalysisRepository,
|
|
||||||
CampaignPeriodComparisonRepository,
|
|
||||||
CampaignRepository,
|
|
||||||
)
|
|
||||||
from agenteval.utils.llm import parse_json_from_llm_text
|
from agenteval.utils.llm import parse_json_from_llm_text
|
||||||
|
|
||||||
_SAME_MOMENT_EPS = 1e-3
|
_SAME_MOMENT_EPS = 1e-3
|
||||||
@ -177,9 +172,7 @@ def validate_comparison_request(
|
|||||||
else:
|
else:
|
||||||
baseline = resolve_auto_baseline(campaign, session)
|
baseline = resolve_auto_baseline(campaign, session)
|
||||||
if baseline is None:
|
if baseline is None:
|
||||||
raise ComparisonError(
|
raise ComparisonError("未找到自动基线:历史活动中没有同计划指纹且已完成分析的活动,可手动选择基线活动")
|
||||||
"未找到自动基线:历史活动中没有同计划指纹且已完成分析的活动,可手动选择基线活动"
|
|
||||||
)
|
|
||||||
|
|
||||||
baseline_analysis = CampaignAnalysisRepository(session).get_by_campaign(baseline.id)
|
baseline_analysis = CampaignAnalysisRepository(session).get_by_campaign(baseline.id)
|
||||||
if baseline_analysis is None or baseline_analysis.status != "completed":
|
if baseline_analysis is None or baseline_analysis.status != "completed":
|
||||||
@ -188,98 +181,6 @@ def validate_comparison_request(
|
|||||||
return baseline
|
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 调用编排 ─────────────────────────────────────────
|
# ── 叙述半边:单次 LLM 调用编排 ─────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -1,26 +1,40 @@
|
|||||||
"""Durable runtime for Campaign intelligence jobs.
|
"""Durable runtime for Campaign intelligence jobs.
|
||||||
|
|
||||||
智能分析与周期对比是两个领域工作 adapter;本 module 统一掌握它们的
|
智能分析与周期对比是两个领域工作 adapter;本 module 的 ``execute``
|
||||||
持久排队、进程内幂等启动、重启恢复和关闭顺序。数据库行是耐久权威,
|
统一掌握它们的结算契约:建行/认领/校验/generating/failed/completed。
|
||||||
TaskRegistry 只保存当前进程中的任务句柄。
|
数据库行是耐久权威,TaskRegistry 只保存当前进程中的任务句柄。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from collections.abc import Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
from sqlmodel import Session
|
from sqlmodel import Session
|
||||||
|
|
||||||
|
from agenteval.storage.async_job_repository import CampaignAnalysisRepository, CampaignPeriodComparisonRepository
|
||||||
from agenteval.storage.db import get_session
|
from agenteval.storage.db import get_session
|
||||||
from agenteval.storage.repository import (
|
|
||||||
CampaignAnalysisRepository,
|
|
||||||
CampaignPeriodComparisonRepository,
|
|
||||||
)
|
|
||||||
from agenteval.task_registry import TaskRegistry
|
from agenteval.task_registry import TaskRegistry
|
||||||
|
|
||||||
_registry = TaskRegistry()
|
_registry = TaskRegistry()
|
||||||
_logger = logging.getLogger("agenteval")
|
_logger = logging.getLogger("agenteval")
|
||||||
MAX_QUEUED_RECOVERY_ATTEMPTS = 3
|
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:
|
def _job_key(kind: str, campaign_id: str) -> str:
|
||||||
@ -50,6 +64,81 @@ def _launch_comparison(
|
|||||||
_registry.launch(_job_key("comparison", campaign_id), run)
|
_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(
|
async def execute_campaign_analysis_job(
|
||||||
campaign_id: str,
|
campaign_id: str,
|
||||||
*,
|
*,
|
||||||
@ -57,53 +146,38 @@ async def execute_campaign_analysis_job(
|
|||||||
chat_client: Any = None,
|
chat_client: Any = None,
|
||||||
session_factory: Optional[Callable[[], Session]] = None,
|
session_factory: Optional[Callable[[], Session]] = None,
|
||||||
) -> 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 (
|
from agenteval.evaluation.analysis import (
|
||||||
analyze_campaign,
|
analyze_campaign,
|
||||||
collect_failure_samples,
|
collect_failure_samples,
|
||||||
gateway_chat_client,
|
gateway_chat_client,
|
||||||
resolve_analysis_model,
|
resolve_analysis_model,
|
||||||
)
|
)
|
||||||
from agenteval.evaluation.comparison import resolve_auto_baseline
|
|
||||||
from agenteval.evaluation.report import load_campaign_view
|
from agenteval.evaluation.report import load_campaign_view
|
||||||
from agenteval.storage.repository import CampaignRepository, RunRepository
|
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)
|
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)
|
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)
|
runs = RunRepository(session).list_by_campaign(campaign_id)
|
||||||
view = load_campaign_view(session, campaign)
|
view = load_campaign_view(session, campaign)
|
||||||
result = await analyze_campaign(
|
result = await analyze_campaign(
|
||||||
@ -111,29 +185,51 @@ async def execute_campaign_analysis_job(
|
|||||||
report=view["report"],
|
report=view["report"],
|
||||||
failure_samples=collect_failure_samples(campaign_id, session),
|
failure_samples=collect_failure_samples(campaign_id, session),
|
||||||
valid_run_ids={run.id for run in runs if run.id},
|
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"],
|
exploration_summary=view["exploration"],
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
return result, {}
|
||||||
_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,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
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:
|
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
|
return
|
||||||
baseline = resolve_auto_baseline(campaign, session)
|
baseline = resolve_auto_baseline(campaign, session)
|
||||||
if baseline is None:
|
if baseline is None:
|
||||||
@ -145,7 +241,7 @@ async def execute_campaign_analysis_job(
|
|||||||
session=session,
|
session=session,
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
_logger.warning("活动 %s 自动周期对比跳过: %s", campaign.id, exc)
|
_logger.warning("活动 %s 自动周期对比跳过: %s", campaign_id, exc)
|
||||||
finally:
|
finally:
|
||||||
session.close()
|
session.close()
|
||||||
|
|
||||||
@ -158,26 +254,20 @@ async def execute_campaign_comparison_job(
|
|||||||
chat_client: Any = None,
|
chat_client: Any = None,
|
||||||
session_factory: Optional[Callable[[], Session]] = None,
|
session_factory: Optional[Callable[[], Session]] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Claim and settle one period-comparison job."""
|
"""Claim and settle one period-comparison job (work adapter)."""
|
||||||
from agenteval.evaluation.analysis import gateway_chat_client, resolve_analysis_model
|
from agenteval.evaluation.comparison import ComparisonError
|
||||||
from agenteval.evaluation.comparison import (
|
|
||||||
ComparisonError,
|
def ensure_queued(session: Session):
|
||||||
compute_metric_diff,
|
from agenteval.evaluation.comparison import validate_comparison_request
|
||||||
narrate_period_comparison,
|
|
||||||
validate_comparison_request,
|
|
||||||
)
|
|
||||||
from agenteval.evaluation.report import load_campaign_report
|
|
||||||
from agenteval.storage.repository import CampaignRepository
|
from agenteval.storage.repository import CampaignRepository
|
||||||
|
|
||||||
session = (session_factory or get_session)()
|
|
||||||
try:
|
|
||||||
comparisons = CampaignPeriodComparisonRepository(session)
|
comparisons = CampaignPeriodComparisonRepository(session)
|
||||||
campaign = CampaignRepository(session).get(campaign_id)
|
campaign = CampaignRepository(session).get(campaign_id)
|
||||||
if campaign is None:
|
if campaign is None:
|
||||||
return
|
return None
|
||||||
|
|
||||||
row = comparisons.get_by_campaign(campaign_id)
|
row = comparisons.get_by_campaign(campaign_id)
|
||||||
if row is None:
|
if row is not None:
|
||||||
|
return row
|
||||||
try:
|
try:
|
||||||
initial_baseline = validate_comparison_request(
|
initial_baseline = validate_comparison_request(
|
||||||
session,
|
session,
|
||||||
@ -191,17 +281,24 @@ async def execute_campaign_comparison_job(
|
|||||||
triggered_by=triggered_by,
|
triggered_by=triggered_by,
|
||||||
error=str(exc),
|
error=str(exc),
|
||||||
)
|
)
|
||||||
return
|
return None
|
||||||
row = comparisons.enqueue(
|
return comparisons.enqueue(
|
||||||
campaign_id,
|
campaign_id,
|
||||||
baseline_campaign_id=initial_baseline.id,
|
baseline_campaign_id=initial_baseline.id,
|
||||||
triggered_by=triggered_by,
|
triggered_by=triggered_by,
|
||||||
)
|
)
|
||||||
if not comparisons.claim_queued(campaign_id).claimed:
|
|
||||||
return
|
|
||||||
|
|
||||||
effective_trigger = row.triggered_by or triggered_by
|
def validate(session: Session) -> dict:
|
||||||
effective_baseline_id = baseline_campaign_id or row.baseline_campaign_id
|
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:
|
try:
|
||||||
baseline = validate_comparison_request(
|
baseline = validate_comparison_request(
|
||||||
session,
|
session,
|
||||||
@ -209,26 +306,23 @@ async def execute_campaign_comparison_job(
|
|||||||
explicit_baseline_id=effective_baseline_id,
|
explicit_baseline_id=effective_baseline_id,
|
||||||
)
|
)
|
||||||
except ComparisonError as exc:
|
except ComparisonError as exc:
|
||||||
comparisons.upsert(
|
raise JobValidationError(str(exc), baseline_campaign_id=effective_baseline_id) from exc
|
||||||
campaign_id,
|
runtime = resolve_analysis_model(campaign, session)
|
||||||
status="failed",
|
return {"baseline_campaign_id": baseline.id, "model_config_id": runtime.id}
|
||||||
baseline_campaign_id=effective_baseline_id,
|
|
||||||
triggered_by=effective_trigger,
|
|
||||||
error=str(exc),
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
|
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)
|
runtime = resolve_analysis_model(campaign, session)
|
||||||
baseline_analysis = CampaignAnalysisRepository(session).get_by_campaign(baseline.id)
|
baseline_analysis = CampaignAnalysisRepository(session).get_by_campaign(baseline.id)
|
||||||
current_analysis = CampaignAnalysisRepository(session).get_by_campaign(campaign_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(
|
diff = compute_metric_diff(
|
||||||
load_campaign_report(session, baseline),
|
load_campaign_report(session, baseline),
|
||||||
load_campaign_report(session, campaign),
|
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"]},
|
valid_scenario_ids={item["scenario_id"] for item in diff["scenarios"]},
|
||||||
chat_client=chat_client or gateway_chat_client(runtime),
|
chat_client=chat_client or gateway_chat_client(runtime),
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
return result, {}
|
||||||
_logger.warning("活动 %s 周期对比失败: %s", campaign_id, exc)
|
|
||||||
comparisons.upsert(
|
await execute(
|
||||||
|
"comparison",
|
||||||
campaign_id,
|
campaign_id,
|
||||||
status="failed",
|
triggered_by=triggered_by,
|
||||||
baseline_campaign_id=baseline.id,
|
repo_cls=CampaignPeriodComparisonRepository,
|
||||||
model_config_id=runtime.id,
|
ensure_queued=ensure_queued,
|
||||||
error=str(exc)[:500],
|
validate=validate,
|
||||||
triggered_by=effective_trigger,
|
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(
|
def enqueue_campaign_analysis(
|
||||||
|
|||||||
@ -22,11 +22,8 @@ from agenteval.exploration.models import ExplorationMessage, ExplorationSession
|
|||||||
from agenteval.models import Campaign
|
from agenteval.models import Campaign
|
||||||
from agenteval.services.model_configs import ModelRuntimeConfig
|
from agenteval.services.model_configs import ModelRuntimeConfig
|
||||||
from agenteval.storage.db import get_session, iso_utc, utc_now
|
from agenteval.storage.db import get_session, iso_utc, utc_now
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.exploration_repository import ExplorationMessageRepository, ExplorationSessionRepository
|
||||||
CampaignRepository,
|
from agenteval.storage.repository import CampaignRepository
|
||||||
ExplorationMessageRepository,
|
|
||||||
ExplorationSessionRepository,
|
|
||||||
)
|
|
||||||
from agenteval.task_registry import TaskRegistry
|
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
|
||||||
|
|
||||||
@ -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)]
|
return [rounds[round(i * span / (limit - 1))] for i in range(limit)]
|
||||||
|
|
||||||
|
|
||||||
def build_judge_messages(
|
def build_judge_messages(session_obj: ExplorationSession, samples: list[ExplorationMessage]) -> list[dict[str, str]]:
|
||||||
session_obj: ExplorationSession, samples: list[ExplorationMessage]
|
|
||||||
) -> list[dict[str, str]]:
|
|
||||||
transcript = []
|
transcript = []
|
||||||
for message in samples:
|
for message in samples:
|
||||||
speaker = "虚拟用户" if message.role == "user" else "被评对象"
|
speaker = "虚拟用户" if message.role == "user" else "被评对象"
|
||||||
|
|||||||
@ -30,10 +30,9 @@ from agenteval.exploration.models import (
|
|||||||
)
|
)
|
||||||
from agenteval.models import CampaignStatus
|
from agenteval.models import CampaignStatus
|
||||||
from agenteval.storage.db import as_utc, utc_now
|
from agenteval.storage.db import as_utc, utc_now
|
||||||
|
from agenteval.storage.exploration_repository import ExplorationMessageRepository, ExplorationSessionRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
ExplorationMessageRepository,
|
|
||||||
ExplorationSessionRepository,
|
|
||||||
TargetRepository,
|
TargetRepository,
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -144,9 +143,7 @@ async def conduct_turn(db_session: Session, *, session_id: str, content: str) ->
|
|||||||
|
|
||||||
received_at = utc_now()
|
received_at = utc_now()
|
||||||
latency_ms = (
|
latency_ms = (
|
||||||
outcome.latency_ms
|
outcome.latency_ms if outcome.latency_ms is not None else int((received_at - sent_at).total_seconds() * 1000)
|
||||||
if outcome.latency_ms is not None
|
|
||||||
else int((received_at - sent_at).total_seconds() * 1000)
|
|
||||||
)
|
)
|
||||||
reply_text = outcome.reply_text or ""
|
reply_text = outcome.reply_text or ""
|
||||||
message_repo.save_message(
|
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.exploration.models import resolve_budget
|
||||||
from agenteval.models import Campaign, CampaignStatus, EvalRun
|
from agenteval.models import Campaign, CampaignStatus, EvalRun
|
||||||
from agenteval.storage.db import as_utc, iso_utc, utc_now
|
from agenteval.storage.db import as_utc, iso_utc, utc_now
|
||||||
|
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
ExplorationSessionRepository,
|
|
||||||
RunRepository,
|
RunRepository,
|
||||||
ScenarioRepository,
|
ScenarioRepository,
|
||||||
TargetRepository,
|
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]]:
|
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))
|
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 sqlalchemy.exc import IntegrityError
|
||||||
from sqlmodel import Session, select
|
from sqlmodel import Session, select
|
||||||
|
|
||||||
from agenteval.exploration.models import ExplorationMessage, ExplorationSession, ExplorationSessionStatus
|
from agenteval.exploration.models import ExplorationSessionStatus
|
||||||
from agenteval.models import (
|
from agenteval.models import (
|
||||||
Campaign,
|
Campaign,
|
||||||
CampaignStatus,
|
CampaignStatus,
|
||||||
@ -26,13 +26,10 @@ from agenteval.models import (
|
|||||||
)
|
)
|
||||||
from agenteval.services.model_configs import ModelConfigService
|
from agenteval.services.model_configs import ModelConfigService
|
||||||
from agenteval.storage.db import (
|
from agenteval.storage.db import (
|
||||||
CampaignAnalysisDB,
|
|
||||||
CampaignDB,
|
CampaignDB,
|
||||||
CampaignPeriodComparisonDB,
|
|
||||||
EvalResultDB,
|
EvalResultDB,
|
||||||
EvalRunDB,
|
EvalRunDB,
|
||||||
EvalTargetDB,
|
EvalTargetDB,
|
||||||
ExplorationMessageDB,
|
|
||||||
ExplorationSessionDB,
|
ExplorationSessionDB,
|
||||||
ScenarioDB,
|
ScenarioDB,
|
||||||
TurnDB,
|
TurnDB,
|
||||||
@ -46,22 +43,6 @@ M = TypeVar("M") # domain model
|
|||||||
DB = TypeVar("DB") # persisted table row
|
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]):
|
class BaseRepository(Generic[M, DB]):
|
||||||
"""Shared CRUD skeleton for id-keyed entity repositories.
|
"""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.EXISTING, self._from_db(existing))
|
||||||
return CampaignRunClaimResult(CampaignRunClaimStatus.CONFLICT)
|
return CampaignRunClaimResult(CampaignRunClaimStatus.CONFLICT)
|
||||||
|
|
||||||
def _get_campaign_identity(
|
def _get_campaign_identity(self, campaign_id: str, plan_index: int, occurrence_index: int) -> Optional[EvalRunDB]:
|
||||||
self, campaign_id: str, plan_index: int, occurrence_index: int
|
|
||||||
) -> Optional[EvalRunDB]:
|
|
||||||
statement = select(EvalRunDB).where(
|
statement = select(EvalRunDB).where(
|
||||||
EvalRunDB.campaign_id == campaign_id,
|
EvalRunDB.campaign_id == campaign_id,
|
||||||
EvalRunDB.campaign_plan_index == plan_index,
|
EvalRunDB.campaign_plan_index == plan_index,
|
||||||
@ -477,11 +456,7 @@ class RunRepository(BaseRepository[EvalRun, EvalRunDB]):
|
|||||||
return len(failed)
|
return len(failed)
|
||||||
|
|
||||||
def _is_recoverable_campaign_pending(self, db: EvalRunDB) -> bool:
|
def _is_recoverable_campaign_pending(self, db: EvalRunDB) -> bool:
|
||||||
if (
|
if db.campaign_id is None or db.campaign_plan_index is None or db.campaign_occurrence_index is None:
|
||||||
db.campaign_id is None
|
|
||||||
or db.campaign_plan_index is None
|
|
||||||
or db.campaign_occurrence_index is None
|
|
||||||
):
|
|
||||||
return False
|
return False
|
||||||
campaign = self.session.get(CampaignDB, db.campaign_id)
|
campaign = self.session.get(CampaignDB, db.campaign_id)
|
||||||
return campaign is not None and campaign.status == CampaignStatus.RUNNING.value
|
return campaign is not None and campaign.status == CampaignStatus.RUNNING.value
|
||||||
@ -779,221 +754,6 @@ class CampaignRepository(BaseRepository[Campaign, CampaignDB]):
|
|||||||
self.session.refresh(db)
|
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:
|
class ResultRepository:
|
||||||
"""Repository for evaluation results."""
|
"""Repository for evaluation results."""
|
||||||
|
|
||||||
@ -1052,129 +812,3 @@ class ResultRepository:
|
|||||||
self.session.commit()
|
self.session.commit()
|
||||||
self.session.refresh(db)
|
self.session.refresh(db)
|
||||||
return _result_from_db(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 pydantic import BaseModel, Field
|
||||||
from sqlmodel import Session
|
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 CampaignCreateError, CampaignLifecycleError
|
||||||
from agenteval.evaluation.campaign_lifecycle import cancel_campaign as cancel_campaign_lifecycle
|
from agenteval.evaluation.campaign_lifecycle import cancel_campaign as cancel_campaign_lifecycle
|
||||||
from agenteval.evaluation.campaign_lifecycle import create_campaign as create_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.comparison import ComparisonError, validate_comparison_request
|
||||||
from agenteval.evaluation.intelligence_jobs import enqueue_campaign_analysis, enqueue_campaign_comparison
|
from agenteval.evaluation.intelligence_jobs import enqueue_campaign_analysis, enqueue_campaign_comparison
|
||||||
from agenteval.evaluation.report_render import render_campaign_markdown
|
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.storage.repository import CampaignRepository
|
||||||
from agenteval.web.deps import get_db
|
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)
|
campaign = CampaignRepository(session).get(campaign_id)
|
||||||
if not campaign:
|
if not campaign:
|
||||||
raise HTTPException(status_code=404, detail="campaign not found")
|
raise HTTPException(status_code=404, detail="campaign not found")
|
||||||
if campaign.status in (CampaignStatus.PLANNED, CampaignStatus.RUNNING):
|
try:
|
||||||
raise HTTPException(status_code=400, detail="活动完成后才能生成智能分析")
|
validate_analysis_request(session, campaign)
|
||||||
if resolve_analysis_model(campaign, session) is None:
|
except AnalysisError as exc:
|
||||||
raise HTTPException(
|
raise HTTPException(status_code=400, detail=str(exc))
|
||||||
status_code=400,
|
|
||||||
detail="未配置分析模型:请在模型配置中心将某个 chat 配置设为「分析默认」,或为该活动指定分析模型",
|
|
||||||
)
|
|
||||||
enqueue_campaign_analysis(campaign_id, triggered_by="manual")
|
enqueue_campaign_analysis(campaign_id, triggered_by="manual")
|
||||||
return {"status": "generating"}
|
return {"status": "generating"}
|
||||||
|
|
||||||
|
|||||||
@ -19,11 +19,8 @@ from agenteval.exploration.errors import (
|
|||||||
ExplorationNotFoundError,
|
ExplorationNotFoundError,
|
||||||
)
|
)
|
||||||
from agenteval.exploration.models import ExplorationTrigger
|
from agenteval.exploration.models import ExplorationTrigger
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.exploration_repository import ExplorationMessageRepository, ExplorationSessionRepository
|
||||||
CampaignRepository,
|
from agenteval.storage.repository import CampaignRepository
|
||||||
ExplorationMessageRepository,
|
|
||||||
ExplorationSessionRepository,
|
|
||||||
)
|
|
||||||
from agenteval.web.deps import get_db
|
from agenteval.web.deps import get_db
|
||||||
|
|
||||||
router = APIRouter()
|
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()
|
@pytest.fixture()
|
||||||
def analysis_spy(monkeypatch):
|
def analysis_spy(monkeypatch):
|
||||||
"""Spy the analysis seam: resolvable model, recorded enqueue calls."""
|
"""Spy the analysis seam: resolvable model, recorded enqueue calls."""
|
||||||
|
from agenteval.evaluation import analysis as analysis_module
|
||||||
|
|
||||||
calls: list[tuple[str, str]] = []
|
calls: list[tuple[str, str]] = []
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
campaign_runner,
|
campaign_runner,
|
||||||
"enqueue_campaign_analysis",
|
"enqueue_campaign_analysis",
|
||||||
lambda cid, *, triggered_by: calls.append((cid, triggered_by)),
|
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
|
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):
|
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)
|
campaign = _make_campaign(seeded_db)
|
||||||
assert runtime.start(campaign.id)
|
assert runtime.start(campaign.id)
|
||||||
await _await_terminal(seeded_db, 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:
|
def _seed_exploration_session(seeded_db, campaign_id: str) -> None:
|
||||||
from agenteval.exploration.models import ExplorationSession
|
from agenteval.exploration.models import ExplorationSession
|
||||||
from agenteval.storage.repository import ExplorationSessionRepository
|
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||||
|
|
||||||
repo = ExplorationSessionRepository(seeded_db)
|
repo = ExplorationSessionRepository(seeded_db)
|
||||||
session_obj = repo.create(
|
session_obj = repo.create(
|
||||||
|
|||||||
@ -11,7 +11,8 @@ from datetime import timedelta
|
|||||||
import pytest
|
import pytest
|
||||||
from agenteval.models import Campaign, ChannelType, EvalTarget, PlatformType, TargetStatus
|
from agenteval.models import Campaign, ChannelType, EvalTarget, PlatformType, TargetStatus
|
||||||
from agenteval.storage.db import ExplorationSessionDB, utc_now
|
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 agenteval.web.app import app
|
||||||
from httpx import ASGITransport, AsyncClient
|
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):
|
async def test_patrol_budget_reflects_existing_sessions(client, seeded_db):
|
||||||
from agenteval.exploration.models import ExplorationSession
|
from agenteval.exploration.models import ExplorationSession
|
||||||
from agenteval.storage.repository import ExplorationSessionRepository
|
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||||
|
|
||||||
ExplorationSessionRepository(seeded_db).create(ExplorationSession(
|
ExplorationSessionRepository(seeded_db).create(ExplorationSession(
|
||||||
campaign_id="c-prod", target_id="t-1", goal="查询账单", persona={"name": "x"},
|
campaign_id="c-prod", target_id="t-1", goal="查询账单", persona={"name": "x"},
|
||||||
|
|||||||
@ -23,9 +23,9 @@ from agenteval.models import (
|
|||||||
TargetStatus,
|
TargetStatus,
|
||||||
)
|
)
|
||||||
from agenteval.storage.db import utc_now
|
from agenteval.storage.db import utc_now
|
||||||
|
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
ExplorationSessionRepository,
|
|
||||||
ScenarioRepository,
|
ScenarioRepository,
|
||||||
TargetRepository,
|
TargetRepository,
|
||||||
)
|
)
|
||||||
|
|||||||
@ -10,7 +10,7 @@ from agenteval.evaluation.analysis import (
|
|||||||
resolve_analysis_model,
|
resolve_analysis_model,
|
||||||
)
|
)
|
||||||
from agenteval.evaluation.intelligence_jobs import execute_campaign_analysis_job
|
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.db import CampaignAnalysisDB, EvalResultDB, ModelConfigDB, TurnDB
|
||||||
from agenteval.storage.model_config_repository import ModelConfigRepository
|
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||||
from agenteval.storage.repository import CampaignRepository, RunRepository
|
from agenteval.storage.repository import CampaignRepository, RunRepository
|
||||||
@ -41,6 +41,7 @@ def _campaign(**overrides) -> Campaign:
|
|||||||
"target_id": "t-1",
|
"target_id": "t-1",
|
||||||
"window_seconds": 86400,
|
"window_seconds": 86400,
|
||||||
"time_scale": 1.0,
|
"time_scale": 1.0,
|
||||||
|
"status": CampaignStatus.COMPLETED,
|
||||||
"plan": [CampaignPlanEntry(scenario_id="s-1", offset_seconds=0, count=2)],
|
"plan": [CampaignPlanEntry(scenario_id="s-1", offset_seconds=0, count=2)],
|
||||||
}
|
}
|
||||||
data.update(overrides)
|
data.update(overrides)
|
||||||
|
|||||||
@ -4,7 +4,8 @@ import pytest
|
|||||||
from agenteval.evaluation.campaign_lifecycle import CampaignLifecycleError, cancel_campaign
|
from agenteval.evaluation.campaign_lifecycle import CampaignLifecycleError, cancel_campaign
|
||||||
from agenteval.exploration.models import ExplorationSession, ExplorationSessionStatus
|
from agenteval.exploration.models import ExplorationSession, ExplorationSessionStatus
|
||||||
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus
|
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
|
from sqlalchemy import event
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -4,7 +4,8 @@ import pytest
|
|||||||
from agenteval.evaluation.campaign_lifecycle import CampaignLifecycleError, complete_campaign
|
from agenteval.evaluation.campaign_lifecycle import CampaignLifecycleError, complete_campaign
|
||||||
from agenteval.exploration.models import ExplorationSession, ExplorationSessionStatus
|
from agenteval.exploration.models import ExplorationSession, ExplorationSessionStatus
|
||||||
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus
|
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
|
from sqlalchemy import event
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -12,8 +12,8 @@ from agenteval.models import (
|
|||||||
EvalTarget,
|
EvalTarget,
|
||||||
Scenario,
|
Scenario,
|
||||||
)
|
)
|
||||||
|
from agenteval.storage.async_job_repository import CampaignAnalysisRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignAnalysisRepository,
|
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
ScenarioRepository,
|
ScenarioRepository,
|
||||||
TargetRepository,
|
TargetRepository,
|
||||||
@ -86,7 +86,7 @@ def test_view_with_comparison(db_session):
|
|||||||
comparisons.upsert(baseline.id, status="completed", result={"summary": "baseline"})
|
comparisons.upsert(baseline.id, status="completed", result={"summary": "baseline"})
|
||||||
comparisons.upsert(campaign.id, status="completed", result={"summary": "current"})
|
comparisons.upsert(campaign.id, status="completed", result={"summary": "current"})
|
||||||
# Create a comparison row
|
# 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 = CampaignPeriodComparisonRepository(db_session)
|
||||||
cmp_repo.upsert(
|
cmp_repo.upsert(
|
||||||
|
|||||||
@ -9,9 +9,9 @@ from datetime import datetime, timezone
|
|||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from agenteval.evaluation.campaign_read_model import load_comparison_view
|
||||||
from agenteval.evaluation.comparison import (
|
from agenteval.evaluation.comparison import (
|
||||||
ComparisonError,
|
ComparisonError,
|
||||||
load_comparison_view,
|
|
||||||
validate_comparison_request,
|
validate_comparison_request,
|
||||||
)
|
)
|
||||||
from agenteval.models import (
|
from agenteval.models import (
|
||||||
@ -23,9 +23,8 @@ from agenteval.models import (
|
|||||||
EvalTarget,
|
EvalTarget,
|
||||||
Scenario,
|
Scenario,
|
||||||
)
|
)
|
||||||
|
from agenteval.storage.async_job_repository import CampaignAnalysisRepository, CampaignPeriodComparisonRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignAnalysisRepository,
|
|
||||||
CampaignPeriodComparisonRepository,
|
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
ScenarioRepository,
|
ScenarioRepository,
|
||||||
TargetRepository,
|
TargetRepository,
|
||||||
|
|||||||
@ -23,10 +23,9 @@ from agenteval.exploration.models import (
|
|||||||
from agenteval.models import Campaign, CampaignPlanEntry
|
from agenteval.models import Campaign, CampaignPlanEntry
|
||||||
from agenteval.storage.db import ModelConfigDB
|
from agenteval.storage.db import ModelConfigDB
|
||||||
from agenteval.storage.model_config_repository import ModelConfigRepository
|
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||||
|
from agenteval.storage.exploration_repository import ExplorationMessageRepository, ExplorationSessionRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
ExplorationMessageRepository,
|
|
||||||
ExplorationSessionRepository,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -22,9 +22,9 @@ from agenteval.models import (
|
|||||||
ExplorationSeeds,
|
ExplorationSeeds,
|
||||||
RunStatus,
|
RunStatus,
|
||||||
)
|
)
|
||||||
|
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
ExplorationSessionRepository,
|
|
||||||
RunRepository,
|
RunRepository,
|
||||||
TargetRepository,
|
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
|
import pytest
|
||||||
from agenteval.evaluation import intelligence_jobs
|
from agenteval.evaluation import intelligence_jobs
|
||||||
from agenteval.exploration import judge
|
from agenteval.exploration import judge
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.async_job_repository import (
|
||||||
AsyncJobClaimStatus,
|
AsyncJobClaimStatus,
|
||||||
CampaignAnalysisRepository,
|
CampaignAnalysisRepository,
|
||||||
CampaignPeriodComparisonRepository,
|
CampaignPeriodComparisonRepository,
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user