| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404 |
- """从 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_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
- 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)
- 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,
- }
- 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) -> str:
- """构建传给 find_agent 的用户消息。"""
- videos_payload = [_serialize_video(video) for video in ctx.videos]
- primary = ctx.primary_video
- seed_video_title = primary.title if primary else ""
- seed_video_id = primary.video_id if primary else ""
- return (
- f"run_id:{run_id}\n"
- f"demand_grade_id:{ctx.demand_grade_id}\n"
- f"demand_word:{ctx.demand_name}\n"
- f"seed_video_id:{seed_video_id}\n"
- f"seed_video_title:{seed_video_title}\n"
- f"reference_videos:{json.dumps(videos_payload, ensure_ascii=False)}\n"
- "说明:video_discovery_run 已由系统预创建,run_id 见上。\n"
- f"第一步必须调用 create_video_discovery_run,并原样传入 run_id={run_id},"
- "同时传入 demand_word、seed_video_title、seed_video_id、demand_grade_id;"
- "relevant_points 请从 reference_videos 中各视频的 points 展平得到。\n"
- "禁止省略 run_id,禁止自行生成新的 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")
- user_input = build_find_agent_user_input(ctx, run_id)
- try:
- agent_result = run_find_agent(user_input)
- return FindDemandExecutionResult(run_id=run_id, agent_result=agent_result)
- 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],
- }
|