batch_save_demand_expansions.py 7.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222
  1. """
  2. 批量保存视频点位拓展判断结果。
  3. """
  4. from __future__ import annotations
  5. import logging
  6. import re
  7. import uuid
  8. from typing import Any
  9. from supply_agent.tools import tool
  10. from supply_infra.db.models.multi_demand_video_point import POINT_TYPES
  11. from supply_infra.db.repositories.demand_video_expansion_repo import (
  12. DemandVideoExpansionRepository,
  13. )
  14. from supply_infra.db.repositories.multi_demand_video_point_repo import (
  15. MultiDemandVideoPointRepository,
  16. )
  17. from supply_infra.db.session import get_session
  18. logger = logging.getLogger(__name__)
  19. _VALID_GRADES = frozenset({"S", "A"})
  20. def _optional_str(value: Any) -> str | None:
  21. if value is None:
  22. return None
  23. text = str(value).strip()
  24. return text or None
  25. def _normalize_items(
  26. items: list[dict[str, Any]],
  27. *,
  28. default_run_id: str,
  29. biz_dt: str,
  30. source_demand_grade_id: int,
  31. source_demand_name: str,
  32. source_grade: str,
  33. ) -> tuple[list[dict[str, Any]], list[str]]:
  34. rows: list[dict[str, Any]] = []
  35. errors: list[str] = []
  36. seen_keys: set[tuple[str, str]] = set()
  37. for idx, item in enumerate(items):
  38. if not isinstance(item, dict):
  39. errors.append(f"第 {idx} 项不是对象")
  40. continue
  41. expanded_text = _optional_str(item.get("expanded_text"))
  42. point_type = _optional_str(item.get("point_type"))
  43. video_id = _optional_str(item.get("video_id"))
  44. reason = _optional_str(item.get("reason"))
  45. if not expanded_text:
  46. errors.append(f"第 {idx} 项缺少 expanded_text")
  47. continue
  48. if not point_type or point_type not in POINT_TYPES:
  49. allowed = " / ".join(POINT_TYPES)
  50. errors.append(f"第 {idx} 项 point_type 无效(只能是 {allowed}): {point_type!r}")
  51. continue
  52. if not video_id:
  53. errors.append(f"第 {idx} 项缺少 video_id")
  54. continue
  55. if not reason:
  56. errors.append(f"第 {idx} 项缺少 reason")
  57. continue
  58. dedupe_key = (expanded_text, video_id)
  59. if dedupe_key in seen_keys:
  60. errors.append(f"第 {idx} 项在本次请求中重复: {expanded_text!r} / {video_id!r}")
  61. continue
  62. seen_keys.add(dedupe_key)
  63. rows.append(
  64. {
  65. "biz_dt": biz_dt,
  66. "run_id": _optional_str(item.get("run_id")) or default_run_id,
  67. "source_demand_grade_id": source_demand_grade_id,
  68. "source_demand_name": source_demand_name,
  69. "source_grade": source_grade,
  70. "expanded_text": expanded_text,
  71. "point_type": point_type,
  72. "point_desc": _optional_str(item.get("point_desc")),
  73. "video_id": video_id,
  74. "reason": reason,
  75. "is_delete": 0,
  76. }
  77. )
  78. return rows, errors
  79. def _normalize_demand_name(name: str) -> str:
  80. text = re.sub(r"\s+", "", name.strip().lower())
  81. return re.sub(r"[,,。..!!??;;::""''\"'、/\\|·—_()()【】\[\]《》<>]", "", text)
  82. def _fill_missing_point_descs(rows: list[dict[str, Any]], session) -> None:
  83. """point_desc 为空时,从 multi_demand_video_point 按 video_id/point_type/point_data 补全。
  84. 查不到或源表 point_desc 也为空时,保持原值(仍为 None)。
  85. """
  86. missing = [row for row in rows if not row.get("point_desc")]
  87. if not missing:
  88. return
  89. video_ids = sorted({str(row["video_id"]) for row in missing})
  90. points_by_vid = MultiDemandVideoPointRepository(session).list_by_video_ids(video_ids)
  91. lookup: dict[tuple[str, str, str], str] = {}
  92. for vid, points in points_by_vid.items():
  93. for point in points:
  94. point_data = _optional_str(point.get("point_data"))
  95. point_type = _optional_str(point.get("point_type"))
  96. point_desc = _optional_str(point.get("point_desc"))
  97. if not point_data or not point_type or not point_desc:
  98. continue
  99. lookup[(vid, point_type, point_data)] = point_desc
  100. for row in missing:
  101. key = (str(row["video_id"]), str(row["point_type"]), str(row["expanded_text"]))
  102. desc = lookup.get(key)
  103. if desc:
  104. row["point_desc"] = desc
  105. # 未命中或源表无描述:不改动,保留空值
  106. @tool
  107. def batch_save_demand_expansions(
  108. items: list[dict[str, Any]],
  109. biz_dt: str,
  110. source_demand_grade_id: int,
  111. source_demand_name: str,
  112. source_grade: str,
  113. run_id: str | None = None,
  114. ) -> str:
  115. """
  116. 保存视频点位拓展判断结果到 demand_video_expansion 表。
  117. Args:
  118. items: 拓展候选列表。每项必填:
  119. - expanded_text: 拓展需求文本(来自 point_data)
  120. - point_type: inspiration / purpose / key
  121. - video_id: 来源视频 id
  122. - reason: 为何与原需求相近、可作为拓展
  123. 选填:
  124. - point_desc: 点位描述快照
  125. biz_dt: 业务日 YYYYMMDD。
  126. source_demand_grade_id: 来源 demand_grade.id(由用户消息提供,原样传入)。
  127. source_demand_name: 来源需求名。
  128. source_grade: 来源等级 S 或 A。
  129. run_id: 任务批次 id;省略则自动生成。
  130. Returns:
  131. 保存结果摘要。若无合适拓展,传 items=[] 即可。
  132. """
  133. biz_dt_text = _optional_str(biz_dt)
  134. if not biz_dt_text or len(biz_dt_text) != 8 or not biz_dt_text.isdigit():
  135. return f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}"
  136. demand_name = _optional_str(source_demand_name)
  137. if not demand_name:
  138. return "source_demand_name 不能为空"
  139. grade = _optional_str(source_grade)
  140. if grade not in _VALID_GRADES:
  141. return f"source_grade 无效,只能是 S 或 A: {source_grade!r}"
  142. try:
  143. grade_id = int(source_demand_grade_id)
  144. except (TypeError, ValueError):
  145. return f"source_demand_grade_id 无效: {source_demand_grade_id!r}"
  146. if not items:
  147. return "无拓展候选,跳过落库"
  148. default_run_id = _optional_str(run_id) or uuid.uuid4().hex
  149. rows, errors = _normalize_items(
  150. items,
  151. default_run_id=default_run_id,
  152. biz_dt=biz_dt_text,
  153. source_demand_grade_id=grade_id,
  154. source_demand_name=demand_name,
  155. source_grade=grade,
  156. )
  157. normalized_demand = _normalize_demand_name(demand_name)
  158. filtered_rows: list[dict[str, Any]] = []
  159. skipped_same: list[str] = []
  160. for row in rows:
  161. if _normalize_demand_name(row["expanded_text"]) == normalized_demand:
  162. skipped_same.append(row["expanded_text"])
  163. continue
  164. filtered_rows.append(row)
  165. if not filtered_rows:
  166. detail = ";".join(errors) if errors else "无有效数据"
  167. same_note = ""
  168. if skipped_same:
  169. same_note = f";剔除与原需求相同 {len(skipped_same)} 条"
  170. return f"没有可保存的数据: {detail}{same_note}"
  171. try:
  172. with get_session() as session:
  173. _fill_missing_point_descs(filtered_rows, session)
  174. saved = DemandVideoExpansionRepository(session).bulk_upsert(filtered_rows)
  175. parts = [f"成功保存 {saved} 条拓展需求", f"biz_dt={biz_dt_text}", f"run_id={default_run_id}"]
  176. if skipped_same:
  177. parts.append(f"剔除与原需求相同 {len(skipped_same)} 条")
  178. if errors:
  179. parts.append(f"校验失败 {len(errors)} 条: " + ";".join(errors[:10]))
  180. message = "。".join(parts)
  181. logger.info("batch_save_demand_expansions completed: %s", message)
  182. return message
  183. except Exception as e:
  184. logger.error("batch_save_demand_expansions failed: %s", e, exc_info=True)
  185. return f"保存拓展需求失败: {e}"