expand_demand_from_video_points.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357
  1. """
  2. 从 S/A 级需求关联视频中挖掘拓展需求。
  3. 任务层负责查库与组装上下文;Agent 仅做语义判断与落库。
  4. """
  5. from __future__ import annotations
  6. import json
  7. import logging
  8. import uuid
  9. from concurrent.futures import ThreadPoolExecutor, as_completed
  10. from datetime import datetime
  11. from typing import Any
  12. from zoneinfo import ZoneInfo
  13. from agents.demand_video_expand_agent.run import (
  14. DemandExpandContext,
  15. VideoPoint,
  16. extract_saved_count,
  17. judge_demand_expansion,
  18. )
  19. from supply_infra.config import get_infra_settings
  20. from supply_infra.db.models.demand_grade import DemandGrade
  21. from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
  22. from supply_infra.db.repositories.demand_video_expansion_repo import (
  23. DemandVideoExpansionRunRepository,
  24. )
  25. from supply_infra.db.repositories.multi_demand_video_point_repo import (
  26. MultiDemandVideoPointRepository,
  27. )
  28. from supply_infra.db.session import get_session
  29. logger = logging.getLogger(__name__)
  30. _DEFAULT_WORKERS = 5
  31. def _resolve_biz_dt(biz_dt: str | None) -> str:
  32. if biz_dt:
  33. text = str(biz_dt).strip()
  34. if len(text) == 8 and text.isdigit():
  35. return text
  36. raise ValueError(f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}")
  37. with get_session() as session:
  38. latest = DemandGradeRepository(session).get_latest_biz_dt()
  39. if latest:
  40. return str(latest)
  41. timezone = ZoneInfo(get_infra_settings().scheduler_timezone)
  42. return datetime.now(timezone).strftime("%Y%m%d")
  43. def _parse_video_ids(raw: Any) -> list[str]:
  44. if raw is None:
  45. return []
  46. items: list[Any]
  47. if isinstance(raw, str):
  48. text = raw.strip()
  49. if not text:
  50. return []
  51. try:
  52. parsed = json.loads(text)
  53. items = list(parsed) if isinstance(parsed, list) else [text]
  54. except (ValueError, TypeError):
  55. items = [part.strip() for part in text.split(",") if part.strip()]
  56. elif isinstance(raw, (list, tuple)):
  57. items = list(raw)
  58. else:
  59. return []
  60. out: list[str] = []
  61. seen: set[str] = set()
  62. for item in items:
  63. vid = str(item).strip() if item is not None else ""
  64. if not vid or vid in seen:
  65. continue
  66. seen.add(vid)
  67. out.append(vid)
  68. return out
  69. def _to_float(value: Any) -> float | None:
  70. if value is None:
  71. return None
  72. return float(value)
  73. def load_expand_contexts(
  74. session,
  75. biz_dt: str,
  76. *,
  77. run_id: str,
  78. skip_finished: bool = True,
  79. ) -> tuple[list[DemandExpandContext], dict[str, int]]:
  80. """加载待拓展判断的需求上下文,并返回跳过统计。"""
  81. stats = {
  82. "total_sa": 0,
  83. "skipped_no_video": 0,
  84. "skipped_no_points": 0,
  85. "skipped_already_done": 0,
  86. }
  87. grades = DemandGradeRepository(session).list_by_biz_dt_and_grades(biz_dt, ("S", "A"))
  88. stats["total_sa"] = len(grades)
  89. finished_ids: set[int] = set()
  90. if skip_finished:
  91. finished_ids = DemandVideoExpansionRunRepository(session).list_finished_grade_ids(
  92. biz_dt
  93. )
  94. rows_with_video: list[tuple[DemandGrade, list[str]]] = []
  95. all_video_ids: set[str] = set()
  96. for row in grades:
  97. if int(row.id) in finished_ids:
  98. stats["skipped_already_done"] += 1
  99. continue
  100. video_ids = _parse_video_ids(row.video_list)
  101. if not video_ids:
  102. stats["skipped_no_video"] += 1
  103. continue
  104. rows_with_video.append((row, video_ids))
  105. all_video_ids.update(video_ids)
  106. points_by_vid = MultiDemandVideoPointRepository(session).list_by_video_ids(all_video_ids)
  107. contexts: list[DemandExpandContext] = []
  108. for row, video_ids in rows_with_video:
  109. points: list[VideoPoint] = []
  110. for vid in video_ids:
  111. for point in points_by_vid.get(vid, []):
  112. points.append(
  113. VideoPoint(
  114. video_id=vid,
  115. point_type=str(point.get("point_type") or ""),
  116. point_data=point.get("point_data"),
  117. point_desc=point.get("point_desc"),
  118. )
  119. )
  120. if not points:
  121. stats["skipped_no_points"] += 1
  122. continue
  123. contexts.append(
  124. DemandExpandContext(
  125. biz_dt=biz_dt,
  126. run_id=run_id,
  127. demand_grade_id=int(row.id),
  128. demand_name=str(row.demand_name),
  129. grade=str(row.grade),
  130. score=_to_float(row.score),
  131. video_ids=video_ids,
  132. points=points,
  133. )
  134. )
  135. return contexts, stats
  136. def _record_run(
  137. *,
  138. biz_dt: str,
  139. run_id: str,
  140. source_demand_grade_id: int,
  141. saved_count: int,
  142. status: str,
  143. error_message: str | None = None,
  144. ) -> None:
  145. with get_session() as session:
  146. DemandVideoExpansionRunRepository(session).upsert_run(
  147. biz_dt=biz_dt,
  148. run_id=run_id,
  149. source_demand_grade_id=source_demand_grade_id,
  150. saved_count=saved_count,
  151. status=status,
  152. error_message=error_message,
  153. )
  154. def _process_single_expand(
  155. ctx: DemandExpandContext,
  156. *,
  157. biz_dt: str,
  158. run_id: str,
  159. ) -> dict[str, Any]:
  160. """并发 worker:对单个需求执行拓展判断并落库执行记录。"""
  161. try:
  162. agent_result = judge_demand_expansion(ctx)
  163. saved_count = extract_saved_count(agent_result)
  164. _record_run(
  165. biz_dt=biz_dt,
  166. run_id=run_id,
  167. source_demand_grade_id=ctx.demand_grade_id,
  168. saved_count=saved_count,
  169. status="finished",
  170. )
  171. logger.info(
  172. "expand demand done: grade_id=%s demand=%s saved=%d iterations=%d",
  173. ctx.demand_grade_id,
  174. ctx.demand_name,
  175. saved_count,
  176. agent_result.iterations,
  177. )
  178. return {
  179. "success": True,
  180. "demand_grade_id": ctx.demand_grade_id,
  181. "demand_name": ctx.demand_name,
  182. "saved_count": saved_count,
  183. "iterations": agent_result.iterations,
  184. }
  185. except Exception as exc:
  186. error_text = str(exc)
  187. _record_run(
  188. biz_dt=biz_dt,
  189. run_id=run_id,
  190. source_demand_grade_id=ctx.demand_grade_id,
  191. saved_count=0,
  192. status="failed",
  193. error_message=error_text,
  194. )
  195. logger.exception(
  196. "expand demand failed: grade_id=%s demand=%s",
  197. ctx.demand_grade_id,
  198. ctx.demand_name,
  199. )
  200. return {
  201. "success": False,
  202. "demand_grade_id": ctx.demand_grade_id,
  203. "demand_name": ctx.demand_name,
  204. "error": error_text,
  205. }
  206. def expand_demand_from_video_points(
  207. biz_dt: str | None = None,
  208. *,
  209. skip_finished: bool = True,
  210. workers: int = _DEFAULT_WORKERS,
  211. ) -> dict[str, Any]:
  212. """
  213. 对指定业务日的 S/A 需求执行视频点位拓展判断。
  214. 程序负责查需求与点位;无视频或无点位则跳过;有数据则并发调用 Agent。
  215. """
  216. started_at = datetime.now()
  217. run_id = uuid.uuid4().hex
  218. try:
  219. resolved_biz_dt = _resolve_biz_dt(biz_dt)
  220. except Exception as exc:
  221. logger.exception("expand_demand_from_video_points preflight failed")
  222. return {
  223. "success": False,
  224. "error": str(exc),
  225. "started_at": started_at.isoformat(),
  226. "finished_at": datetime.now().isoformat(),
  227. }
  228. with get_session() as session:
  229. contexts, preload_stats = load_expand_contexts(
  230. session,
  231. resolved_biz_dt,
  232. run_id=run_id,
  233. skip_finished=skip_finished,
  234. )
  235. logger.info(
  236. "expand_demand_from_video_points start: biz_dt=%s run_id=%s workers=%s pending=%s",
  237. resolved_biz_dt,
  238. run_id,
  239. workers,
  240. len(contexts),
  241. )
  242. result: dict[str, Any] = {
  243. "success": True,
  244. "run_id": run_id,
  245. "biz_dt": resolved_biz_dt,
  246. "started_at": started_at.isoformat(),
  247. "workers": 0,
  248. **preload_stats,
  249. "processed": 0,
  250. "saved_total": 0,
  251. "failed": 0,
  252. "errors": [],
  253. }
  254. if not contexts:
  255. finished_at = datetime.now()
  256. result["finished_at"] = finished_at.isoformat()
  257. result["duration_seconds"] = round((finished_at - started_at).total_seconds(), 2)
  258. logger.info("expand_demand_from_video_points finished: %s", result)
  259. return result
  260. worker_count = max(1, min(int(workers), len(contexts)))
  261. result["workers"] = worker_count
  262. with ThreadPoolExecutor(max_workers=worker_count) as executor:
  263. futures = [
  264. executor.submit(_process_single_expand, ctx, biz_dt=resolved_biz_dt, run_id=run_id)
  265. for ctx in contexts
  266. ]
  267. for future in as_completed(futures):
  268. try:
  269. item_result = future.result()
  270. except Exception as exc:
  271. logger.exception("expand demand worker 出现未捕获错误: biz_dt=%s", resolved_biz_dt)
  272. result["failed"] += 1
  273. result["errors"].append({"error": str(exc)})
  274. continue
  275. result["processed"] += 1
  276. if item_result.get("success"):
  277. result["saved_total"] += int(item_result.get("saved_count") or 0)
  278. continue
  279. result["failed"] += 1
  280. result["errors"].append(
  281. {
  282. "demand_grade_id": item_result.get("demand_grade_id"),
  283. "demand_name": item_result.get("demand_name"),
  284. "error": item_result.get("error"),
  285. }
  286. )
  287. finished_at = datetime.now()
  288. result["finished_at"] = finished_at.isoformat()
  289. result["duration_seconds"] = round((finished_at - started_at).total_seconds(), 2)
  290. result["success"] = result["failed"] == 0
  291. logger.info("expand_demand_from_video_points finished: %s", result)
  292. return result
  293. if __name__ == "__main__":
  294. import sys
  295. from supply_infra.scheduler.cli_result import run_cli
  296. _args = sys.argv[1:]
  297. _biz_dt = _args[0] if _args and not _args[0].startswith("-") else None
  298. _workers_arg = None
  299. if _biz_dt and len(_args) > 1 and not _args[1].startswith("-"):
  300. _workers_arg = _args[1]
  301. _skip_finished = "--force" not in _args
  302. run_cli(
  303. lambda: expand_demand_from_video_points(
  304. _biz_dt,
  305. skip_finished=_skip_finished,
  306. workers=int(_workers_arg) if _workers_arg else 5,
  307. ),
  308. label="expand_demand_from_video_points",
  309. )