demand_run.py 13 KB

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