"""从 demand_grade + demand_video_expansion 组装 find_agent 输入并执行。""" from __future__ import annotations import json import logging import uuid from collections.abc import Iterable from dataclasses import dataclass, field from datetime import datetime from typing import Any from zoneinfo import ZoneInfo from supply_infra.video_discovery_gates import build_rule_snapshot from supply_agent.types import AgentResult from supply_infra.config import get_infra_settings from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository from supply_infra.db.repositories.demand_video_expansion_repo import ( DemandVideoExpansionRepository, ) from supply_infra.db.repositories.multi_demand_video_detail_repo import ( MultiDemandVideoDetailRepository, ) from supply_infra.db.session import get_session from supply_infra.services.video_discovery_service import get_video_discovery_service logger = logging.getLogger(__name__) _POINT_TYPES = {"inspiration", "purpose", "key"} @dataclass class FindDemandPoint: point: str point_type: str point_desc: str | None = None @dataclass class FindDemandVideo: video_id: str title: str points: list[FindDemandPoint] = field(default_factory=list) @dataclass class FindDemandContext: """单条 find_agent 任务:一个 S/A 需求词及其下全部视频与全部拓展点位。""" biz_dt: str demand_grade_id: int demand_name: str grade: str videos: list[FindDemandVideo] = field(default_factory=list) @property def video_count(self) -> int: return len(self.videos) @property def point_count(self) -> int: return sum(len(video.points) for video in self.videos) @property def primary_video(self) -> FindDemandVideo | None: return self.videos[0] if self.videos else None @dataclass class FindDemandExecutionResult: """单条需求找视频执行结果。""" run_id: str | None = None skipped: bool = False skip_reason: str | None = None agent_result: AgentResult | None = None succeeded: bool = False failure_reason: str | None = None def _resolve_biz_dt(biz_dt: str | None) -> str: if biz_dt: text = str(biz_dt).strip() if len(text) == 8 and text.isdigit(): return text raise ValueError(f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}") with get_session() as session: latest = DemandGradeRepository(session).get_latest_biz_dt() if latest: return str(latest) timezone = ZoneInfo(get_infra_settings().scheduler_timezone) return datetime.now(timezone).strftime("%Y%m%d") def _build_video_points(expansions: list[Any]) -> list[FindDemandPoint]: points: list[FindDemandPoint] = [] seen: set[tuple[str, str]] = set() for row in expansions: point_type = str(row.point_type or "").strip() if point_type not in _POINT_TYPES: continue expanded_text = str(row.expanded_text or "").strip() if not expanded_text: continue dedupe_key = (point_type, expanded_text) if dedupe_key in seen: continue seen.add(dedupe_key) point_desc = str(row.point_desc or "").strip() or None points.append( FindDemandPoint( point=expanded_text, point_type=point_type, point_desc=point_desc, ) ) return points def _serialize_video(video: FindDemandVideo) -> dict[str, Any]: return { "video_id": video.video_id, "title": video.title, "points": [ { "point": point.point, "point_type": point.point_type, **({"point_desc": point.point_desc} if point.point_desc else {}), } for point in video.points ], } def _grade_priority_key(row: Any) -> tuple[float, int, str, int]: """S 优先于 A;同等级按 score 降序。""" score = float(row.score) if row.score is not None else -1.0 grade_rank = 0 if str(row.grade) == "S" else 1 return (-score, grade_rank, str(row.demand_name), int(row.id)) def load_find_demand_contexts( session, biz_dt: str, *, grades: Iterable[str] = ("S", "A"), top_limit: int | None = None, ) -> list[FindDemandContext]: """加载指定业务日全部 S/A 需求(可选 top_limit 截断),每个需求词组装为一条完整上下文。""" grades_rows = DemandGradeRepository(session).list_by_biz_dt_and_grades(biz_dt, grades) if not grades_rows: return [] grades_rows = sorted(grades_rows, key=_grade_priority_key) if top_limit is not None and top_limit > 0: grades_rows = grades_rows[: int(top_limit)] expansion_repo = DemandVideoExpansionRepository(session) detail_repo = MultiDemandVideoDetailRepository(session) contexts: list[FindDemandContext] = [] for grade_row in grades_rows: expansions = expansion_repo.list_by_demand_grade(biz_dt, int(grade_row.id)) if not expansions: continue by_video: dict[str, list[Any]] = {} video_order: list[str] = [] for row in expansions: video_id = str(row.video_id or "").strip() if not video_id: continue if video_id not in by_video: by_video[video_id] = [] video_order.append(video_id) by_video[video_id].append(row) if not video_order: continue details = detail_repo.list_by_vids(video_order) videos: list[FindDemandVideo] = [] for video_id in video_order: points = _build_video_points(by_video[video_id]) if not points: continue detail = details.get(video_id) title = str(detail.title).strip() if detail and detail.title else "" videos.append( FindDemandVideo( video_id=video_id, title=title or f"(无标题|{video_id})", points=points, ) ) if not videos: continue contexts.append( FindDemandContext( biz_dt=biz_dt, demand_grade_id=int(grade_row.id), demand_name=str(grade_row.demand_name), grade=str(grade_row.grade), videos=videos, ) ) return contexts def pick_find_demand_context( biz_dt: str | None = None, *, index: int = 0, demand_grade_id: int | None = None, grades: Iterable[str] = ("S", "A"), ) -> FindDemandContext | None: """从数据库选取一条待执行的 find_agent 上下文。""" resolved_biz_dt = _resolve_biz_dt(biz_dt) with get_session() as session: contexts = load_find_demand_contexts( session, resolved_biz_dt, grades=grades, ) if demand_grade_id is not None: for ctx in contexts: if ctx.demand_grade_id == int(demand_grade_id): return ctx return None if 0 <= index < len(contexts): return contexts[index] return None def list_find_demand_contexts( biz_dt: str | None = None, *, grades: Iterable[str] = ("S", "A"), top_limit: int | None = None, ) -> tuple[str, list[FindDemandContext]]: """返回解析后的业务日与待执行上下文(默认当日全部 S/A,可选 top_limit)。""" resolved_biz_dt = _resolve_biz_dt(biz_dt) with get_session() as session: contexts = load_find_demand_contexts( session, resolved_biz_dt, grades=grades, top_limit=top_limit, ) return resolved_biz_dt, contexts def build_run_input_payload(ctx: FindDemandContext) -> dict[str, Any]: """构建写入 video_discovery_run 的输入快照。""" return { "biz_dt": ctx.biz_dt, "demand_grade_id": ctx.demand_grade_id, "demand_name": ctx.demand_name, "grade": ctx.grade, "relevant_points": _flatten_relevant_points(ctx), "reference_videos": [_serialize_video(video) for video in ctx.videos], } def prepare_video_discovery_run( ctx: FindDemandContext, *, force: bool = False, ) -> tuple[str | None, str | None]: """执行 Agent 前预创建 video_discovery_run,返回 (run_id, skip_reason)。""" payload = build_run_input_payload(ctx) rule_snapshot = build_rule_snapshot() primary = ctx.primary_video values = { "run_id": uuid.uuid4().hex, "biz_dt": ctx.biz_dt, "demand_grade_id": ctx.demand_grade_id, "demand_word": ctx.demand_name, "seed_video_id": primary.video_id if primary else None, "seed_video_title": primary.title if primary else None, "relevant_points_json": json.dumps(payload, ensure_ascii=False), "intent_summary": None, "status": "running", "stop_reason": None, "rule_version": rule_snapshot["rule_version"], "rule_config_json": json.dumps(rule_snapshot, ensure_ascii=False), } run_id, skip_reason = get_video_discovery_service().prepare_scheduled_run( biz_dt=ctx.biz_dt, demand_grade_id=ctx.demand_grade_id, values=values, force=force, ) if skip_reason is None and run_id: logger.info( "prepared video_discovery_run: biz_dt=%s demand_grade_id=%s run_id=%s demand=%s", ctx.biz_dt, ctx.demand_grade_id, run_id, ctx.demand_name, ) return run_id, skip_reason def filter_pending_contexts( contexts: list[FindDemandContext], biz_dt: str, *, skip_finished: bool = True, ) -> tuple[list[FindDemandContext], dict[str, int]]: """过滤当天已执行过的需求,返回待执行列表与跳过统计。""" stats = {"skipped_already_done": 0} if not skip_finished or not contexts: return contexts, stats skip_ids = get_video_discovery_service().list_skip_grade_ids(biz_dt) pending = [ctx for ctx in contexts if ctx.demand_grade_id not in skip_ids] stats["skipped_already_done"] = len(contexts) - len(pending) return pending, stats def _flatten_relevant_points(ctx: FindDemandContext) -> list[dict[str, Any]]: """将全部视频点位拍平,并保留来源视频信息。""" relevant_points: list[dict[str, Any]] = [] for video in ctx.videos: for point in video.points: item: dict[str, Any] = { "point": point.point, "point_type": point.point_type, "video_id": video.video_id, "video_title": video.title, } if point.point_desc: item["point_desc"] = point.point_desc relevant_points.append(item) return relevant_points def build_find_agent_user_input( ctx: FindDemandContext, run_id: str, *, rule_snapshot: dict[str, Any] | None = None, ) -> str: """构建传给 find_agent 的用户消息。""" videos_payload = [_serialize_video(video) for video in ctx.videos] active_rules = rule_snapshot or build_rule_snapshot() return ( f"run_id:{run_id}\n" f"demand_grade_id:{ctx.demand_grade_id}\n" f"demand_word:{ctx.demand_name}\n" f"current_datetime:{active_rules['current_datetime']}\n" f"current_date:{active_rules['current_date']}\n" f"timezone:{active_rules['timezone']}\n" f"quality_gate_rules:{json.dumps(active_rules, ensure_ascii=False)}\n" f"reference_videos:{json.dumps(videos_payload, ensure_ascii=False)}\n" "说明:video_discovery_run 已由系统预创建,无需也不得由模型再次创建。\n" "请直接依据 reference_videos 中各视频的 title 和 points 理解需求并开始搜索;" f"后续所有存储和状态工具都必须使用预创建的 run_id={run_id}。" ) def discover_videos_for_demand( ctx: FindDemandContext, *, force: bool = False, ) -> FindDemandExecutionResult: """对单条需求记录执行 find_agent。""" from agents.find_agent import run_find_agent run_id, skip_reason = prepare_video_discovery_run(ctx, force=force) if skip_reason: logger.info( "skip find_agent: demand_grade_id=%s demand=%s reason=%s", ctx.demand_grade_id, ctx.demand_name, skip_reason, ) return FindDemandExecutionResult(skipped=True, skip_reason=skip_reason) if not run_id: raise RuntimeError("prepare_video_discovery_run 未返回 run_id") run_snapshot = get_video_discovery_service().lookup_run(run_id) or {} user_input = build_find_agent_user_input( ctx, run_id, rule_snapshot=run_snapshot.get("rule_config"), ) try: from agents.find_agent.run_outcome import evaluate_find_agent_run agent_result = run_find_agent(user_input, run_id=run_id) outcome = evaluate_find_agent_run(run_id, agent_result) if not outcome.succeeded: stop_reason = outcome.failure_reason or "incomplete" content_preview = (agent_result.content or "").strip()[:500] if content_preview: stop_reason = f"{stop_reason}: {content_preview}" get_video_discovery_service().mark_run_failed( run_id, stop_reason=stop_reason, ) return FindDemandExecutionResult( run_id=run_id, agent_result=agent_result, succeeded=outcome.succeeded, failure_reason=outcome.failure_reason, ) except Exception as exc: get_video_discovery_service().mark_run_failed( run_id, stop_reason=str(exc), ) raise def serialize_find_demand_context(ctx: FindDemandContext) -> dict[str, Any]: """便于日志/测试脚本输出的结构化摘要。""" return { "biz_dt": ctx.biz_dt, "demand_grade_id": ctx.demand_grade_id, "demand_name": ctx.demand_name, "grade": ctx.grade, "video_count": ctx.video_count, "point_count": ctx.point_count, "videos": [_serialize_video(video) for video in ctx.videos], }