demand_run.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435
  1. """从 demand_grade + demand_video_expansion 组装 find_agent 输入并执行。"""
  2. from __future__ import annotations
  3. import json
  4. import logging
  5. import uuid
  6. from collections.abc import Iterable
  7. from dataclasses import dataclass, field
  8. from datetime import datetime
  9. from typing import Any
  10. from zoneinfo import ZoneInfo
  11. from supply_infra.video_discovery_gates import build_rule_snapshot
  12. from supply_agent.types import AgentResult
  13. from supply_infra.config import get_infra_settings
  14. from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
  15. from supply_infra.db.repositories.demand_video_expansion_repo import (
  16. DemandVideoExpansionRepository,
  17. )
  18. from supply_infra.db.repositories.multi_demand_video_detail_repo import (
  19. MultiDemandVideoDetailRepository,
  20. )
  21. from supply_infra.db.session import get_session
  22. from supply_infra.services.video_discovery_service import get_video_discovery_service
  23. logger = logging.getLogger(__name__)
  24. _POINT_TYPES = {"inspiration", "purpose", "key"}
  25. @dataclass
  26. class FindDemandPoint:
  27. point: str
  28. point_type: str
  29. point_desc: str | None = None
  30. @dataclass
  31. class FindDemandVideo:
  32. video_id: str
  33. title: str
  34. points: list[FindDemandPoint] = field(default_factory=list)
  35. @dataclass
  36. class FindDemandContext:
  37. """单条 find_agent 任务:一个 S/A 需求词及其下全部视频与全部拓展点位。"""
  38. biz_dt: str
  39. demand_grade_id: int
  40. demand_name: str
  41. grade: str
  42. videos: list[FindDemandVideo] = field(default_factory=list)
  43. @property
  44. def video_count(self) -> int:
  45. return len(self.videos)
  46. @property
  47. def point_count(self) -> int:
  48. return sum(len(video.points) for video in self.videos)
  49. @property
  50. def primary_video(self) -> FindDemandVideo | None:
  51. return self.videos[0] if self.videos else None
  52. @dataclass
  53. class FindDemandExecutionResult:
  54. """单条需求找视频执行结果。"""
  55. run_id: str | None = None
  56. skipped: bool = False
  57. skip_reason: str | None = None
  58. agent_result: AgentResult | None = None
  59. succeeded: bool = False
  60. failure_reason: str | None = None
  61. def _resolve_biz_dt(biz_dt: str | None) -> str:
  62. if biz_dt:
  63. text = str(biz_dt).strip()
  64. if len(text) == 8 and text.isdigit():
  65. return text
  66. raise ValueError(f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}")
  67. with get_session() as session:
  68. latest = DemandGradeRepository(session).get_latest_biz_dt()
  69. if latest:
  70. return str(latest)
  71. timezone = ZoneInfo(get_infra_settings().scheduler_timezone)
  72. return datetime.now(timezone).strftime("%Y%m%d")
  73. def _build_video_points(expansions: list[Any]) -> list[FindDemandPoint]:
  74. points: list[FindDemandPoint] = []
  75. seen: set[tuple[str, str]] = set()
  76. for row in expansions:
  77. point_type = str(row.point_type or "").strip()
  78. if point_type not in _POINT_TYPES:
  79. continue
  80. expanded_text = str(row.expanded_text or "").strip()
  81. if not expanded_text:
  82. continue
  83. dedupe_key = (point_type, expanded_text)
  84. if dedupe_key in seen:
  85. continue
  86. seen.add(dedupe_key)
  87. point_desc = str(row.point_desc or "").strip() or None
  88. points.append(
  89. FindDemandPoint(
  90. point=expanded_text,
  91. point_type=point_type,
  92. point_desc=point_desc,
  93. )
  94. )
  95. return points
  96. def _serialize_video(video: FindDemandVideo) -> dict[str, Any]:
  97. return {
  98. "video_id": video.video_id,
  99. "title": video.title,
  100. "points": [
  101. {
  102. "point": point.point,
  103. "point_type": point.point_type,
  104. **({"point_desc": point.point_desc} if point.point_desc else {}),
  105. }
  106. for point in video.points
  107. ],
  108. }
  109. def _grade_priority_key(row: Any) -> tuple[float, int, str, int]:
  110. """S 优先于 A;同等级按 score 降序。"""
  111. score = float(row.score) if row.score is not None else -1.0
  112. grade_rank = 0 if str(row.grade) == "S" else 1
  113. return (-score, grade_rank, str(row.demand_name), int(row.id))
  114. def load_find_demand_contexts(
  115. session,
  116. biz_dt: str,
  117. *,
  118. grades: Iterable[str] = ("S", "A"),
  119. top_limit: int | None = None,
  120. ) -> list[FindDemandContext]:
  121. """加载指定业务日全部 S/A 需求(可选 top_limit 截断),每个需求词组装为一条完整上下文。"""
  122. grades_rows = DemandGradeRepository(session).list_by_biz_dt_and_grades(biz_dt, grades)
  123. if not grades_rows:
  124. return []
  125. grades_rows = sorted(grades_rows, key=_grade_priority_key)
  126. if top_limit is not None and top_limit > 0:
  127. grades_rows = grades_rows[: int(top_limit)]
  128. expansion_repo = DemandVideoExpansionRepository(session)
  129. detail_repo = MultiDemandVideoDetailRepository(session)
  130. contexts: list[FindDemandContext] = []
  131. for grade_row in grades_rows:
  132. expansions = expansion_repo.list_by_demand_grade(biz_dt, int(grade_row.id))
  133. if not expansions:
  134. continue
  135. by_video: dict[str, list[Any]] = {}
  136. video_order: list[str] = []
  137. for row in expansions:
  138. video_id = str(row.video_id or "").strip()
  139. if not video_id:
  140. continue
  141. if video_id not in by_video:
  142. by_video[video_id] = []
  143. video_order.append(video_id)
  144. by_video[video_id].append(row)
  145. if not video_order:
  146. continue
  147. details = detail_repo.list_by_vids(video_order)
  148. videos: list[FindDemandVideo] = []
  149. for video_id in video_order:
  150. points = _build_video_points(by_video[video_id])
  151. if not points:
  152. continue
  153. detail = details.get(video_id)
  154. title = str(detail.title).strip() if detail and detail.title else ""
  155. videos.append(
  156. FindDemandVideo(
  157. video_id=video_id,
  158. title=title or f"(无标题|{video_id})",
  159. points=points,
  160. )
  161. )
  162. if not videos:
  163. continue
  164. contexts.append(
  165. FindDemandContext(
  166. biz_dt=biz_dt,
  167. demand_grade_id=int(grade_row.id),
  168. demand_name=str(grade_row.demand_name),
  169. grade=str(grade_row.grade),
  170. videos=videos,
  171. )
  172. )
  173. return contexts
  174. def pick_find_demand_context(
  175. biz_dt: str | None = None,
  176. *,
  177. index: int = 0,
  178. demand_grade_id: int | None = None,
  179. grades: Iterable[str] = ("S", "A"),
  180. ) -> FindDemandContext | None:
  181. """从数据库选取一条待执行的 find_agent 上下文。"""
  182. resolved_biz_dt = _resolve_biz_dt(biz_dt)
  183. with get_session() as session:
  184. contexts = load_find_demand_contexts(
  185. session,
  186. resolved_biz_dt,
  187. grades=grades,
  188. )
  189. if demand_grade_id is not None:
  190. for ctx in contexts:
  191. if ctx.demand_grade_id == int(demand_grade_id):
  192. return ctx
  193. return None
  194. if 0 <= index < len(contexts):
  195. return contexts[index]
  196. return None
  197. def list_find_demand_contexts(
  198. biz_dt: str | None = None,
  199. *,
  200. grades: Iterable[str] = ("S", "A"),
  201. top_limit: int | None = None,
  202. ) -> tuple[str, list[FindDemandContext]]:
  203. """返回解析后的业务日与待执行上下文(默认当日全部 S/A,可选 top_limit)。"""
  204. resolved_biz_dt = _resolve_biz_dt(biz_dt)
  205. with get_session() as session:
  206. contexts = load_find_demand_contexts(
  207. session,
  208. resolved_biz_dt,
  209. grades=grades,
  210. top_limit=top_limit,
  211. )
  212. return resolved_biz_dt, contexts
  213. def build_run_input_payload(ctx: FindDemandContext) -> dict[str, Any]:
  214. """构建写入 video_discovery_run 的输入快照。"""
  215. return {
  216. "biz_dt": ctx.biz_dt,
  217. "demand_grade_id": ctx.demand_grade_id,
  218. "demand_name": ctx.demand_name,
  219. "grade": ctx.grade,
  220. "relevant_points": _flatten_relevant_points(ctx),
  221. "reference_videos": [_serialize_video(video) for video in ctx.videos],
  222. }
  223. def prepare_video_discovery_run(
  224. ctx: FindDemandContext,
  225. *,
  226. force: bool = False,
  227. ) -> tuple[str | None, str | None]:
  228. """执行 Agent 前预创建 video_discovery_run,返回 (run_id, skip_reason)。"""
  229. payload = build_run_input_payload(ctx)
  230. rule_snapshot = build_rule_snapshot()
  231. primary = ctx.primary_video
  232. values = {
  233. "run_id": uuid.uuid4().hex,
  234. "biz_dt": ctx.biz_dt,
  235. "demand_grade_id": ctx.demand_grade_id,
  236. "demand_word": ctx.demand_name,
  237. "seed_video_id": primary.video_id if primary else None,
  238. "seed_video_title": primary.title if primary else None,
  239. "relevant_points_json": json.dumps(payload, ensure_ascii=False),
  240. "intent_summary": None,
  241. "status": "running",
  242. "stop_reason": None,
  243. "rule_version": rule_snapshot["rule_version"],
  244. "rule_config_json": json.dumps(rule_snapshot, ensure_ascii=False),
  245. }
  246. run_id, skip_reason = get_video_discovery_service().prepare_scheduled_run(
  247. biz_dt=ctx.biz_dt,
  248. demand_grade_id=ctx.demand_grade_id,
  249. values=values,
  250. force=force,
  251. )
  252. if skip_reason is None and run_id:
  253. logger.info(
  254. "prepared video_discovery_run: biz_dt=%s demand_grade_id=%s run_id=%s demand=%s",
  255. ctx.biz_dt,
  256. ctx.demand_grade_id,
  257. run_id,
  258. ctx.demand_name,
  259. )
  260. return run_id, skip_reason
  261. def filter_pending_contexts(
  262. contexts: list[FindDemandContext],
  263. biz_dt: str,
  264. *,
  265. skip_finished: bool = True,
  266. ) -> tuple[list[FindDemandContext], dict[str, int]]:
  267. """过滤当天已执行过的需求,返回待执行列表与跳过统计。"""
  268. stats = {"skipped_already_done": 0}
  269. if not skip_finished or not contexts:
  270. return contexts, stats
  271. skip_ids = get_video_discovery_service().list_skip_grade_ids(biz_dt)
  272. pending = [ctx for ctx in contexts if ctx.demand_grade_id not in skip_ids]
  273. stats["skipped_already_done"] = len(contexts) - len(pending)
  274. return pending, stats
  275. def _flatten_relevant_points(ctx: FindDemandContext) -> list[dict[str, Any]]:
  276. """将全部视频点位拍平,并保留来源视频信息。"""
  277. relevant_points: list[dict[str, Any]] = []
  278. for video in ctx.videos:
  279. for point in video.points:
  280. item: dict[str, Any] = {
  281. "point": point.point,
  282. "point_type": point.point_type,
  283. "video_id": video.video_id,
  284. "video_title": video.title,
  285. }
  286. if point.point_desc:
  287. item["point_desc"] = point.point_desc
  288. relevant_points.append(item)
  289. return relevant_points
  290. def build_find_agent_user_input(
  291. ctx: FindDemandContext,
  292. run_id: str,
  293. *,
  294. rule_snapshot: dict[str, Any] | None = None,
  295. ) -> str:
  296. """构建传给 find_agent 的用户消息。"""
  297. videos_payload = [_serialize_video(video) for video in ctx.videos]
  298. active_rules = rule_snapshot or build_rule_snapshot()
  299. return (
  300. f"run_id:{run_id}\n"
  301. f"demand_grade_id:{ctx.demand_grade_id}\n"
  302. f"demand_word:{ctx.demand_name}\n"
  303. f"current_datetime:{active_rules['current_datetime']}\n"
  304. f"current_date:{active_rules['current_date']}\n"
  305. f"timezone:{active_rules['timezone']}\n"
  306. f"quality_gate_rules:{json.dumps(active_rules, ensure_ascii=False)}\n"
  307. f"reference_videos:{json.dumps(videos_payload, ensure_ascii=False)}\n"
  308. "说明:video_discovery_run 已由系统预创建,无需也不得由模型再次创建。\n"
  309. "请直接依据 reference_videos 中各视频的 title 和 points 理解需求并开始搜索;"
  310. f"后续所有存储和状态工具都必须使用预创建的 run_id={run_id}。"
  311. )
  312. def discover_videos_for_demand(
  313. ctx: FindDemandContext,
  314. *,
  315. force: bool = False,
  316. ) -> FindDemandExecutionResult:
  317. """对单条需求记录执行 find_agent。"""
  318. from agents.find_agent import run_find_agent
  319. run_id, skip_reason = prepare_video_discovery_run(ctx, force=force)
  320. if skip_reason:
  321. logger.info(
  322. "skip find_agent: demand_grade_id=%s demand=%s reason=%s",
  323. ctx.demand_grade_id,
  324. ctx.demand_name,
  325. skip_reason,
  326. )
  327. return FindDemandExecutionResult(skipped=True, skip_reason=skip_reason)
  328. if not run_id:
  329. raise RuntimeError("prepare_video_discovery_run 未返回 run_id")
  330. run_snapshot = get_video_discovery_service().lookup_run(run_id) or {}
  331. user_input = build_find_agent_user_input(
  332. ctx,
  333. run_id,
  334. rule_snapshot=run_snapshot.get("rule_config"),
  335. )
  336. try:
  337. from agents.find_agent.run_outcome import evaluate_find_agent_run
  338. agent_result = run_find_agent(user_input, run_id=run_id)
  339. outcome = evaluate_find_agent_run(run_id, agent_result)
  340. if not outcome.succeeded:
  341. stop_reason = outcome.failure_reason or "incomplete"
  342. content_preview = (agent_result.content or "").strip()[:500]
  343. if content_preview:
  344. stop_reason = f"{stop_reason}: {content_preview}"
  345. get_video_discovery_service().mark_run_failed(
  346. run_id,
  347. stop_reason=stop_reason,
  348. )
  349. return FindDemandExecutionResult(
  350. run_id=run_id,
  351. agent_result=agent_result,
  352. succeeded=outcome.succeeded,
  353. failure_reason=outcome.failure_reason,
  354. )
  355. except Exception as exc:
  356. get_video_discovery_service().mark_run_failed(
  357. run_id,
  358. stop_reason=str(exc),
  359. )
  360. raise
  361. def serialize_find_demand_context(ctx: FindDemandContext) -> dict[str, Any]:
  362. """便于日志/测试脚本输出的结构化摘要。"""
  363. return {
  364. "biz_dt": ctx.biz_dt,
  365. "demand_grade_id": ctx.demand_grade_id,
  366. "demand_name": ctx.demand_name,
  367. "grade": ctx.grade,
  368. "video_count": ctx.video_count,
  369. "point_count": ctx.point_count,
  370. "videos": [_serialize_video(video) for video in ctx.videos],
  371. }