demand_run.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523
  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 sqlalchemy import select
  12. from supply_infra.video_discovery_gates import build_rule_snapshot
  13. from supply_agent.types import AgentResult
  14. from supply_infra.config import get_infra_settings
  15. from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
  16. from supply_infra.db.models.business_harness import (
  17. DailyDemandPackage,
  18. DailyDemandTask,
  19. PlatformDemandVersion,
  20. )
  21. from supply_infra.db.repositories.demand_video_expansion_repo import (
  22. DemandVideoExpansionRepository,
  23. )
  24. from supply_infra.db.repositories.multi_demand_video_detail_repo import (
  25. MultiDemandVideoDetailRepository,
  26. )
  27. from supply_infra.db.session import get_session
  28. from supply_infra.services.video_discovery_service import get_video_discovery_service
  29. logger = logging.getLogger(__name__)
  30. _POINT_TYPES = {"inspiration", "purpose", "key"}
  31. @dataclass
  32. class FindDemandPoint:
  33. point: str
  34. point_type: str
  35. point_desc: str | None = None
  36. @dataclass
  37. class FindDemandVideo:
  38. video_id: str
  39. title: str
  40. points: list[FindDemandPoint] = field(default_factory=list)
  41. @dataclass
  42. class FindDemandContext:
  43. """单条 find_agent 任务:一个 S/A 需求词及其下全部视频与全部拓展点位。"""
  44. biz_dt: str
  45. demand_grade_id: int
  46. demand_name: str
  47. grade: str
  48. platform_demand_id: str | None = None
  49. platform_demand_version_id: str | None = None
  50. demand_package_id: str | None = None
  51. daily_demand_task_id: str | None = None
  52. task_payload: dict[str, Any] | None = None
  53. videos: list[FindDemandVideo] = field(default_factory=list)
  54. @property
  55. def video_count(self) -> int:
  56. return len(self.videos)
  57. @property
  58. def point_count(self) -> int:
  59. return sum(len(video.points) for video in self.videos)
  60. @property
  61. def primary_video(self) -> FindDemandVideo | None:
  62. return self.videos[0] if self.videos else None
  63. @dataclass
  64. class FindDemandExecutionResult:
  65. """单条需求找视频执行结果。"""
  66. run_id: str | None = None
  67. skipped: bool = False
  68. skip_reason: str | None = None
  69. agent_result: AgentResult | None = None
  70. succeeded: bool = False
  71. failure_reason: str | None = None
  72. def _resolve_biz_dt(biz_dt: str | None) -> str:
  73. if biz_dt:
  74. text = str(biz_dt).strip()
  75. if len(text) == 8 and text.isdigit():
  76. return text
  77. raise ValueError(f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}")
  78. with get_session() as session:
  79. latest = DemandGradeRepository(session).get_latest_biz_dt()
  80. if latest:
  81. return str(latest)
  82. timezone = ZoneInfo(get_infra_settings().scheduler_timezone)
  83. return datetime.now(timezone).strftime("%Y%m%d")
  84. def _build_video_points(expansions: list[Any]) -> list[FindDemandPoint]:
  85. points: list[FindDemandPoint] = []
  86. seen: set[tuple[str, str]] = set()
  87. for row in expansions:
  88. point_type = str(row.point_type or "").strip()
  89. if point_type not in _POINT_TYPES:
  90. continue
  91. expanded_text = str(row.expanded_text or "").strip()
  92. if not expanded_text:
  93. continue
  94. dedupe_key = (point_type, expanded_text)
  95. if dedupe_key in seen:
  96. continue
  97. seen.add(dedupe_key)
  98. point_desc = str(row.point_desc or "").strip() or None
  99. points.append(
  100. FindDemandPoint(
  101. point=expanded_text,
  102. point_type=point_type,
  103. point_desc=point_desc,
  104. )
  105. )
  106. return points
  107. def _serialize_video(video: FindDemandVideo) -> dict[str, Any]:
  108. return {
  109. "video_id": video.video_id,
  110. "title": video.title,
  111. "points": [
  112. {
  113. "point": point.point,
  114. "point_type": point.point_type,
  115. **({"point_desc": point.point_desc} if point.point_desc else {}),
  116. }
  117. for point in video.points
  118. ],
  119. }
  120. def _grade_priority_key(row: Any) -> tuple[float, int, str, int]:
  121. """S 优先于 A;同等级按 score 降序。"""
  122. score = float(row.score) if row.score is not None else -1.0
  123. grade_rank = 0 if str(row.grade) == "S" else 1
  124. return (-score, grade_rank, str(row.demand_name), int(row.id))
  125. def load_find_demand_contexts(
  126. session,
  127. biz_dt: str,
  128. *,
  129. grades: Iterable[str] = ("S", "A"),
  130. top_limit: int | None = None,
  131. package_run_id: str | None = None,
  132. ) -> list[FindDemandContext]:
  133. """从唯一已发布任务包加载需求;不再由 Find Agent 自行解释等级。"""
  134. package_stmt = select(DailyDemandPackage).where(
  135. DailyDemandPackage.biz_dt == biz_dt,
  136. DailyDemandPackage.status == "published",
  137. )
  138. if package_run_id:
  139. package_stmt = package_stmt.where(
  140. DailyDemandPackage.run_id == package_run_id
  141. )
  142. package = session.scalar(
  143. package_stmt.order_by(
  144. DailyDemandPackage.package_version.desc(),
  145. DailyDemandPackage.created_at.desc(),
  146. ).limit(1)
  147. )
  148. if package is None:
  149. raise RuntimeError(
  150. f"biz_dt={biz_dt} 没有已发布的统一需求任务包,禁止回退拼接 DemandGrade"
  151. )
  152. task_rows = session.execute(
  153. select(DailyDemandTask, PlatformDemandVersion)
  154. .join(
  155. PlatformDemandVersion,
  156. PlatformDemandVersion.platform_demand_version_id
  157. == DailyDemandTask.platform_demand_version_id,
  158. )
  159. .where(
  160. DailyDemandTask.demand_package_id == package.demand_package_id,
  161. DailyDemandTask.action_tier.in_(
  162. (
  163. "assure_supply",
  164. "priority",
  165. "targeted_validation",
  166. "exploration",
  167. )
  168. ),
  169. )
  170. ).all()
  171. tier_order = {
  172. "assure_supply": 0,
  173. "priority": 1,
  174. "targeted_validation": 2,
  175. "exploration": 3,
  176. }
  177. task_rows = sorted(
  178. task_rows,
  179. key=lambda pair: (
  180. tier_order.get(str(pair[0].action_tier), 99),
  181. -float(pair[0].task_payload_json.get("local_supply_priority") or 0),
  182. str(pair[1].name),
  183. ),
  184. )
  185. if top_limit is not None and top_limit > 0:
  186. task_rows = task_rows[: int(top_limit)]
  187. expansion_repo = DemandVideoExpansionRepository(session)
  188. detail_repo = MultiDemandVideoDetailRepository(session)
  189. contexts: list[FindDemandContext] = []
  190. for task_row, platform_version in task_rows:
  191. grade_row = session.get(
  192. DemandGradeRepository.model,
  193. int(platform_version.source_demand_grade_id),
  194. )
  195. if grade_row is None:
  196. raise RuntimeError(
  197. "统一任务包引用的 DemandGrade 不存在: "
  198. f"{platform_version.source_demand_grade_id}"
  199. )
  200. expansions = expansion_repo.list_by_demand_grade(biz_dt, int(grade_row.id))
  201. by_video: dict[str, list[Any]] = {}
  202. video_order: list[str] = []
  203. for row in expansions:
  204. video_id = str(row.video_id or "").strip()
  205. if not video_id:
  206. continue
  207. if video_id not in by_video:
  208. by_video[video_id] = []
  209. video_order.append(video_id)
  210. by_video[video_id].append(row)
  211. details = detail_repo.list_by_vids(video_order)
  212. videos: list[FindDemandVideo] = []
  213. for video_id in video_order:
  214. points = _build_video_points(by_video[video_id])
  215. if not points:
  216. continue
  217. detail = details.get(video_id)
  218. title = str(detail.title).strip() if detail and detail.title else ""
  219. videos.append(
  220. FindDemandVideo(
  221. video_id=video_id,
  222. title=title or f"(无标题|{video_id})",
  223. points=points,
  224. )
  225. )
  226. contexts.append(
  227. FindDemandContext(
  228. biz_dt=biz_dt,
  229. demand_grade_id=int(grade_row.id),
  230. demand_name=str(grade_row.demand_name),
  231. grade=str(grade_row.grade),
  232. platform_demand_id=platform_version.platform_demand_id,
  233. platform_demand_version_id=(
  234. platform_version.platform_demand_version_id
  235. ),
  236. demand_package_id=package.demand_package_id,
  237. daily_demand_task_id=task_row.task_id,
  238. task_payload=dict(task_row.task_payload_json),
  239. videos=videos,
  240. )
  241. )
  242. return contexts
  243. def pick_find_demand_context(
  244. biz_dt: str | None = None,
  245. *,
  246. index: int = 0,
  247. demand_grade_id: int | None = None,
  248. grades: Iterable[str] = ("S", "A"),
  249. package_run_id: str | None = None,
  250. ) -> FindDemandContext | None:
  251. """从数据库选取一条待执行的 find_agent 上下文。"""
  252. resolved_biz_dt = _resolve_biz_dt(biz_dt)
  253. with get_session() as session:
  254. contexts = load_find_demand_contexts(
  255. session,
  256. resolved_biz_dt,
  257. grades=grades,
  258. package_run_id=package_run_id,
  259. )
  260. if demand_grade_id is not None:
  261. for ctx in contexts:
  262. if ctx.demand_grade_id == int(demand_grade_id):
  263. return ctx
  264. return None
  265. if 0 <= index < len(contexts):
  266. return contexts[index]
  267. return None
  268. def list_find_demand_contexts(
  269. biz_dt: str | None = None,
  270. *,
  271. grades: Iterable[str] = ("S", "A"),
  272. top_limit: int | None = None,
  273. package_run_id: str | None = None,
  274. ) -> tuple[str, list[FindDemandContext]]:
  275. """返回解析后的业务日与待执行上下文(默认当日全部 S/A,可选 top_limit)。"""
  276. resolved_biz_dt = _resolve_biz_dt(biz_dt)
  277. with get_session() as session:
  278. contexts = load_find_demand_contexts(
  279. session,
  280. resolved_biz_dt,
  281. grades=grades,
  282. top_limit=top_limit,
  283. package_run_id=package_run_id,
  284. )
  285. return resolved_biz_dt, contexts
  286. def build_run_input_payload(ctx: FindDemandContext) -> dict[str, Any]:
  287. """构建写入 video_discovery_run 的输入快照。"""
  288. return {
  289. "biz_dt": ctx.biz_dt,
  290. "demand_grade_id": ctx.demand_grade_id,
  291. "demand_name": ctx.demand_name,
  292. "grade": ctx.grade,
  293. "platform_demand_id": ctx.platform_demand_id,
  294. "platform_demand_version_id": ctx.platform_demand_version_id,
  295. "demand_package_id": ctx.demand_package_id,
  296. "daily_demand_task_id": ctx.daily_demand_task_id,
  297. "daily_demand_task": ctx.task_payload,
  298. "relevant_points": _flatten_relevant_points(ctx),
  299. "reference_videos": [_serialize_video(video) for video in ctx.videos],
  300. }
  301. def prepare_video_discovery_run(
  302. ctx: FindDemandContext,
  303. *,
  304. force: bool = False,
  305. ) -> tuple[str | None, str | None]:
  306. """执行 Agent 前预创建 video_discovery_run,返回 (run_id, skip_reason)。"""
  307. payload = build_run_input_payload(ctx)
  308. rule_snapshot = build_rule_snapshot()
  309. primary = ctx.primary_video
  310. values = {
  311. "run_id": uuid.uuid4().hex,
  312. "biz_dt": ctx.biz_dt,
  313. "demand_grade_id": ctx.demand_grade_id,
  314. "demand_package_id": ctx.demand_package_id,
  315. "daily_demand_task_id": ctx.daily_demand_task_id,
  316. "platform_demand_version_id": ctx.platform_demand_version_id,
  317. "demand_word": ctx.demand_name,
  318. "seed_video_id": primary.video_id if primary else None,
  319. "seed_video_title": primary.title if primary else None,
  320. "relevant_points_json": json.dumps(payload, ensure_ascii=False),
  321. "intent_summary": None,
  322. "status": "running",
  323. "stop_reason": None,
  324. "rule_version": rule_snapshot["rule_version"],
  325. "rule_config_json": json.dumps(rule_snapshot, ensure_ascii=False),
  326. }
  327. run_id, skip_reason = get_video_discovery_service().prepare_scheduled_run(
  328. biz_dt=ctx.biz_dt,
  329. demand_grade_id=ctx.demand_grade_id,
  330. values=values,
  331. force=force,
  332. )
  333. if skip_reason is None and run_id:
  334. logger.info(
  335. "prepared video_discovery_run: biz_dt=%s demand_grade_id=%s run_id=%s demand=%s",
  336. ctx.biz_dt,
  337. ctx.demand_grade_id,
  338. run_id,
  339. ctx.demand_name,
  340. )
  341. return run_id, skip_reason
  342. def filter_pending_contexts(
  343. contexts: list[FindDemandContext],
  344. biz_dt: str,
  345. *,
  346. skip_finished: bool = True,
  347. ) -> tuple[list[FindDemandContext], dict[str, int]]:
  348. """过滤当天已执行过的需求,返回待执行列表与跳过统计。"""
  349. stats = {"skipped_already_done": 0}
  350. if not skip_finished or not contexts:
  351. return contexts, stats
  352. skip_ids = get_video_discovery_service().list_skip_grade_ids(biz_dt)
  353. pending = [ctx for ctx in contexts if ctx.demand_grade_id not in skip_ids]
  354. stats["skipped_already_done"] = len(contexts) - len(pending)
  355. return pending, stats
  356. def _flatten_relevant_points(ctx: FindDemandContext) -> list[dict[str, Any]]:
  357. """将全部视频点位拍平,并保留来源视频信息。"""
  358. relevant_points: list[dict[str, Any]] = []
  359. for video in ctx.videos:
  360. for point in video.points:
  361. item: dict[str, Any] = {
  362. "point": point.point,
  363. "point_type": point.point_type,
  364. "video_id": video.video_id,
  365. "video_title": video.title,
  366. }
  367. if point.point_desc:
  368. item["point_desc"] = point.point_desc
  369. relevant_points.append(item)
  370. return relevant_points
  371. def build_find_agent_user_input(
  372. ctx: FindDemandContext,
  373. run_id: str,
  374. *,
  375. rule_snapshot: dict[str, Any] | None = None,
  376. ) -> str:
  377. """构建传给 find_agent 的用户消息。"""
  378. videos_payload = [_serialize_video(video) for video in ctx.videos]
  379. active_rules = rule_snapshot or build_rule_snapshot()
  380. return (
  381. f"run_id:{run_id}\n"
  382. f"demand_grade_id:{ctx.demand_grade_id}\n"
  383. f"demand_word:{ctx.demand_name}\n"
  384. f"platform_demand_id:{ctx.platform_demand_id}\n"
  385. f"platform_demand_version_id:{ctx.platform_demand_version_id}\n"
  386. f"demand_package_id:{ctx.demand_package_id}\n"
  387. f"daily_demand_task:{json.dumps(ctx.task_payload or {}, ensure_ascii=False)}\n"
  388. f"current_datetime:{active_rules['current_datetime']}\n"
  389. f"current_date:{active_rules['current_date']}\n"
  390. f"timezone:{active_rules['timezone']}\n"
  391. f"quality_gate_rules:{json.dumps(active_rules, ensure_ascii=False)}\n"
  392. f"reference_videos:{json.dumps(videos_payload, ensure_ascii=False)}\n"
  393. "说明:video_discovery_run 已由系统预创建,无需也不得由模型再次创建。\n"
  394. "请以 daily_demand_task 为唯一业务契约,结合 reference_videos 理解需求并开始搜索;"
  395. f"后续所有存储和状态工具都必须使用预创建的 run_id={run_id}。"
  396. )
  397. def discover_videos_for_demand(
  398. ctx: FindDemandContext,
  399. *,
  400. force: bool = False,
  401. ) -> FindDemandExecutionResult:
  402. """对单条需求记录执行 find_agent。"""
  403. from agents.find_agent import run_find_agent
  404. run_id, skip_reason = prepare_video_discovery_run(ctx, force=force)
  405. if skip_reason:
  406. logger.info(
  407. "skip find_agent: demand_grade_id=%s demand=%s reason=%s",
  408. ctx.demand_grade_id,
  409. ctx.demand_name,
  410. skip_reason,
  411. )
  412. return FindDemandExecutionResult(skipped=True, skip_reason=skip_reason)
  413. if not run_id:
  414. raise RuntimeError("prepare_video_discovery_run 未返回 run_id")
  415. run_snapshot = get_video_discovery_service().lookup_run(run_id) or {}
  416. user_input = build_find_agent_user_input(
  417. ctx,
  418. run_id,
  419. rule_snapshot=run_snapshot.get("rule_config"),
  420. )
  421. try:
  422. from agents.find_agent.run_outcome import evaluate_find_agent_run
  423. agent_result = run_find_agent(user_input, run_id=run_id)
  424. outcome = evaluate_find_agent_run(run_id, agent_result)
  425. if not outcome.succeeded:
  426. stop_reason = outcome.failure_reason or "incomplete"
  427. content_preview = (agent_result.content or "").strip()[:500]
  428. if content_preview:
  429. stop_reason = f"{stop_reason}: {content_preview}"
  430. get_video_discovery_service().mark_run_failed(
  431. run_id,
  432. stop_reason=stop_reason,
  433. )
  434. return FindDemandExecutionResult(
  435. run_id=run_id,
  436. agent_result=agent_result,
  437. succeeded=outcome.succeeded,
  438. failure_reason=outcome.failure_reason,
  439. )
  440. except Exception as exc:
  441. get_video_discovery_service().mark_run_failed(
  442. run_id,
  443. stop_reason=str(exc),
  444. )
  445. raise
  446. def serialize_find_demand_context(ctx: FindDemandContext) -> dict[str, Any]:
  447. """便于日志/测试脚本输出的结构化摘要。"""
  448. return {
  449. "biz_dt": ctx.biz_dt,
  450. "demand_grade_id": ctx.demand_grade_id,
  451. "demand_name": ctx.demand_name,
  452. "grade": ctx.grade,
  453. "platform_demand_id": ctx.platform_demand_id,
  454. "platform_demand_version_id": ctx.platform_demand_version_id,
  455. "demand_package_id": ctx.demand_package_id,
  456. "daily_demand_task_id": ctx.daily_demand_task_id,
  457. "task_payload": ctx.task_payload,
  458. "video_count": ctx.video_count,
  459. "point_count": ctx.point_count,
  460. "videos": [_serialize_video(video) for video in ctx.videos],
  461. }