348 lines
12 KiB
Python
348 lines
12 KiB
Python
"""Durable runtime for Campaign intelligence jobs.
|
||
|
||
智能分析与周期对比是两个领域工作 adapter;本 module 统一掌握它们的
|
||
持久排队、进程内幂等启动、重启恢复和关闭顺序。数据库行是耐久权威,
|
||
TaskRegistry 只保存当前进程中的任务句柄。
|
||
"""
|
||
|
||
import logging
|
||
from collections.abc import Callable
|
||
from typing import Any, Optional
|
||
|
||
from sqlmodel import Session
|
||
|
||
from agenteval.storage.db import get_session
|
||
from agenteval.storage.repository import (
|
||
CampaignAnalysisRepository,
|
||
CampaignPeriodComparisonRepository,
|
||
)
|
||
from agenteval.task_registry import TaskRegistry
|
||
|
||
_registry = TaskRegistry()
|
||
_logger = logging.getLogger("agenteval")
|
||
MAX_QUEUED_RECOVERY_ATTEMPTS = 3
|
||
|
||
|
||
def _job_key(kind: str, campaign_id: str) -> str:
|
||
return f"{kind}:{campaign_id}"
|
||
|
||
|
||
def _launch_analysis(campaign_id: str, *, triggered_by: str) -> None:
|
||
async def run(_cancel) -> None:
|
||
await execute_campaign_analysis_job(campaign_id, triggered_by=triggered_by)
|
||
|
||
_registry.launch(_job_key("analysis", campaign_id), run)
|
||
|
||
|
||
def _launch_comparison(
|
||
campaign_id: str,
|
||
*,
|
||
triggered_by: str,
|
||
baseline_campaign_id: str,
|
||
) -> None:
|
||
async def run(_cancel) -> None:
|
||
await execute_campaign_comparison_job(
|
||
campaign_id,
|
||
triggered_by=triggered_by,
|
||
baseline_campaign_id=baseline_campaign_id,
|
||
)
|
||
|
||
_registry.launch(_job_key("comparison", campaign_id), run)
|
||
|
||
|
||
async def execute_campaign_analysis_job(
|
||
campaign_id: str,
|
||
*,
|
||
triggered_by: str,
|
||
chat_client: Any = None,
|
||
session_factory: Optional[Callable[[], Session]] = None,
|
||
) -> None:
|
||
"""Claim and settle one intelligent-analysis job."""
|
||
from agenteval.evaluation.analysis import (
|
||
analyze_campaign,
|
||
collect_failure_samples,
|
||
gateway_chat_client,
|
||
resolve_analysis_model,
|
||
)
|
||
from agenteval.evaluation.comparison import resolve_auto_baseline
|
||
from agenteval.evaluation.report import load_campaign_view
|
||
from agenteval.storage.repository import CampaignRepository, RunRepository
|
||
|
||
session = (session_factory or get_session)()
|
||
try:
|
||
analyses = CampaignAnalysisRepository(session)
|
||
row = analyses.get_by_campaign(campaign_id)
|
||
if row is None:
|
||
row = analyses.enqueue(campaign_id, triggered_by=triggered_by)
|
||
if not analyses.claim_queued(campaign_id).claimed:
|
||
return
|
||
effective_trigger = row.triggered_by or triggered_by
|
||
|
||
campaign = CampaignRepository(session).get(campaign_id)
|
||
if campaign is None:
|
||
analyses.upsert(
|
||
campaign_id,
|
||
status="failed",
|
||
triggered_by=effective_trigger,
|
||
error="campaign not found",
|
||
)
|
||
return
|
||
runtime = resolve_analysis_model(campaign, session)
|
||
if runtime is None:
|
||
analyses.upsert(
|
||
campaign_id,
|
||
status="failed",
|
||
triggered_by=effective_trigger,
|
||
error="未配置分析模型:请在模型配置中心将某个 chat 配置设为「分析默认」",
|
||
)
|
||
return
|
||
analyses.upsert(
|
||
campaign_id,
|
||
status="generating",
|
||
model_config_id=runtime.id,
|
||
triggered_by=effective_trigger,
|
||
)
|
||
try:
|
||
client = chat_client or gateway_chat_client(runtime)
|
||
runs = RunRepository(session).list_by_campaign(campaign_id)
|
||
view = load_campaign_view(session, campaign)
|
||
result = await analyze_campaign(
|
||
campaign=campaign,
|
||
report=view["report"],
|
||
failure_samples=collect_failure_samples(campaign_id, session),
|
||
valid_run_ids={run.id for run in runs if run.id},
|
||
chat_client=client,
|
||
exploration_summary=view["exploration"],
|
||
)
|
||
except Exception as exc:
|
||
_logger.warning("活动 %s 智能分析失败: %s", campaign_id, exc)
|
||
analyses.upsert(
|
||
campaign_id,
|
||
status="failed",
|
||
model_config_id=runtime.id,
|
||
error=str(exc)[:500],
|
||
triggered_by=effective_trigger,
|
||
)
|
||
return
|
||
analyses.upsert(
|
||
campaign_id,
|
||
status="completed",
|
||
result=result,
|
||
model_config_id=runtime.id,
|
||
triggered_by=effective_trigger,
|
||
)
|
||
|
||
try:
|
||
if campaign.time_scale != 1 or resolve_analysis_model(campaign, session) is None:
|
||
return
|
||
baseline = resolve_auto_baseline(campaign, session)
|
||
if baseline is None:
|
||
return
|
||
enqueue_campaign_comparison(
|
||
campaign.id,
|
||
triggered_by="auto",
|
||
baseline_campaign_id=baseline.id,
|
||
session=session,
|
||
)
|
||
except Exception as exc:
|
||
_logger.warning("活动 %s 自动周期对比跳过: %s", campaign.id, exc)
|
||
finally:
|
||
session.close()
|
||
|
||
|
||
async def execute_campaign_comparison_job(
|
||
campaign_id: str,
|
||
*,
|
||
triggered_by: str,
|
||
baseline_campaign_id: Optional[str] = None,
|
||
chat_client: Any = None,
|
||
session_factory: Optional[Callable[[], Session]] = None,
|
||
) -> None:
|
||
"""Claim and settle one period-comparison job."""
|
||
from agenteval.evaluation.analysis import gateway_chat_client, resolve_analysis_model
|
||
from agenteval.evaluation.comparison import (
|
||
ComparisonError,
|
||
compute_metric_diff,
|
||
narrate_period_comparison,
|
||
validate_comparison_request,
|
||
)
|
||
from agenteval.evaluation.report import load_campaign_report
|
||
from agenteval.storage.repository import CampaignRepository
|
||
|
||
session = (session_factory or get_session)()
|
||
try:
|
||
comparisons = CampaignPeriodComparisonRepository(session)
|
||
campaign = CampaignRepository(session).get(campaign_id)
|
||
if campaign is None:
|
||
return
|
||
|
||
row = comparisons.get_by_campaign(campaign_id)
|
||
if row is None:
|
||
try:
|
||
initial_baseline = validate_comparison_request(
|
||
session,
|
||
campaign,
|
||
explicit_baseline_id=baseline_campaign_id,
|
||
)
|
||
except ComparisonError as exc:
|
||
comparisons.upsert(
|
||
campaign_id,
|
||
status="failed",
|
||
triggered_by=triggered_by,
|
||
error=str(exc),
|
||
)
|
||
return
|
||
row = comparisons.enqueue(
|
||
campaign_id,
|
||
baseline_campaign_id=initial_baseline.id,
|
||
triggered_by=triggered_by,
|
||
)
|
||
if not comparisons.claim_queued(campaign_id).claimed:
|
||
return
|
||
|
||
effective_trigger = row.triggered_by or triggered_by
|
||
effective_baseline_id = baseline_campaign_id or row.baseline_campaign_id
|
||
try:
|
||
baseline = validate_comparison_request(
|
||
session,
|
||
campaign,
|
||
explicit_baseline_id=effective_baseline_id,
|
||
)
|
||
except ComparisonError as exc:
|
||
comparisons.upsert(
|
||
campaign_id,
|
||
status="failed",
|
||
baseline_campaign_id=effective_baseline_id,
|
||
triggered_by=effective_trigger,
|
||
error=str(exc),
|
||
)
|
||
return
|
||
|
||
runtime = resolve_analysis_model(campaign, session)
|
||
baseline_analysis = CampaignAnalysisRepository(session).get_by_campaign(baseline.id)
|
||
current_analysis = CampaignAnalysisRepository(session).get_by_campaign(campaign_id)
|
||
comparisons.upsert(
|
||
campaign_id,
|
||
status="generating",
|
||
baseline_campaign_id=baseline.id,
|
||
model_config_id=runtime.id,
|
||
triggered_by=effective_trigger,
|
||
)
|
||
try:
|
||
diff = compute_metric_diff(
|
||
load_campaign_report(session, baseline),
|
||
load_campaign_report(session, campaign),
|
||
)
|
||
result = await narrate_period_comparison(
|
||
baseline_analysis=baseline_analysis.get_result(),
|
||
current_analysis=current_analysis.get_result(),
|
||
metric_diff=diff,
|
||
valid_scenario_ids={item["scenario_id"] for item in diff["scenarios"]},
|
||
chat_client=chat_client or gateway_chat_client(runtime),
|
||
)
|
||
except Exception as exc:
|
||
_logger.warning("活动 %s 周期对比失败: %s", campaign_id, exc)
|
||
comparisons.upsert(
|
||
campaign_id,
|
||
status="failed",
|
||
baseline_campaign_id=baseline.id,
|
||
model_config_id=runtime.id,
|
||
error=str(exc)[:500],
|
||
triggered_by=effective_trigger,
|
||
)
|
||
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(
|
||
campaign_id: str,
|
||
*,
|
||
triggered_by: str,
|
||
session: Optional[Session] = None,
|
||
) -> None:
|
||
"""Persist an analysis job, then launch its process-local worker."""
|
||
owns_session = session is None
|
||
active_session = session or get_session()
|
||
try:
|
||
row = CampaignAnalysisRepository(active_session).enqueue(
|
||
campaign_id,
|
||
triggered_by=triggered_by,
|
||
)
|
||
finally:
|
||
if owns_session:
|
||
active_session.close()
|
||
if row.status == "queued":
|
||
_launch_analysis(campaign_id, triggered_by=row.triggered_by or triggered_by)
|
||
|
||
|
||
def enqueue_campaign_comparison(
|
||
campaign_id: str,
|
||
*,
|
||
triggered_by: str,
|
||
baseline_campaign_id: str,
|
||
session: Optional[Session] = None,
|
||
) -> None:
|
||
"""Persist a comparison job, then launch its process-local worker."""
|
||
owns_session = session is None
|
||
active_session = session or get_session()
|
||
try:
|
||
row = CampaignPeriodComparisonRepository(active_session).enqueue(
|
||
campaign_id,
|
||
baseline_campaign_id=baseline_campaign_id,
|
||
triggered_by=triggered_by,
|
||
)
|
||
finally:
|
||
if owns_session:
|
||
active_session.close()
|
||
if row.status == "queued":
|
||
_launch_comparison(
|
||
campaign_id,
|
||
triggered_by=row.triggered_by or triggered_by,
|
||
baseline_campaign_id=row.baseline_campaign_id,
|
||
)
|
||
|
||
|
||
def recover_campaign_intelligence_jobs(session: Session) -> tuple[int, int]:
|
||
"""Fail interrupted work and relaunch every durably queued job."""
|
||
analyses = CampaignAnalysisRepository(session)
|
||
comparisons = CampaignPeriodComparisonRepository(session)
|
||
interrupted = analyses.mark_orphans_failed() + comparisons.mark_orphans_failed()
|
||
|
||
queued_analyses = analyses.prepare_queued_recovery(
|
||
MAX_QUEUED_RECOVERY_ATTEMPTS,
|
||
"服务重启恢复次数超过上限,分析任务已终止",
|
||
)
|
||
for row in queued_analyses:
|
||
_launch_analysis(row.campaign_id, triggered_by=row.triggered_by or "manual")
|
||
|
||
queued_comparisons = comparisons.prepare_queued_recovery(
|
||
MAX_QUEUED_RECOVERY_ATTEMPTS,
|
||
"服务重启恢复次数超过上限,周期对比任务已终止",
|
||
)
|
||
for row in queued_comparisons:
|
||
_launch_comparison(
|
||
row.campaign_id,
|
||
triggered_by=row.triggered_by or "manual",
|
||
baseline_campaign_id=row.baseline_campaign_id,
|
||
)
|
||
return interrupted, len(queued_analyses) + len(queued_comparisons)
|
||
|
||
|
||
def is_intelligence_job_running(kind: str, campaign_id: str) -> bool:
|
||
"""Expose process-local liveness without exposing registry internals."""
|
||
return _registry.is_running(_job_key(kind, campaign_id))
|
||
|
||
|
||
async def shutdown_campaign_intelligence_jobs() -> None:
|
||
"""Stop every live analysis and comparison worker."""
|
||
await _registry.shutdown_all()
|