batch_save_demand_grades.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253
  1. """
  2. 批量保存需求分级结果到 demand_grade 表。
  3. """
  4. from __future__ import annotations
  5. import logging
  6. from decimal import Decimal
  7. from typing import Any, Optional
  8. from agents.demand_grade_agent.tools.shared import (
  9. VALID_GRADES,
  10. collect_strategies,
  11. dump_int_list,
  12. merge_video_ids,
  13. normalize_biz_dt,
  14. )
  15. from supply_agent.tools import tool
  16. from supply_infra.db.repositories.demand_grade_category_rel_repo import (
  17. DemandGradeCategoryRelRepository,
  18. )
  19. from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
  20. from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
  21. from supply_infra.db.session import get_session
  22. logger = logging.getLogger(__name__)
  23. def _optional_decimal(value: Any, field: str, idx: int, errors: list[str]) -> Decimal | None:
  24. if value is None or value == "":
  25. return None
  26. try:
  27. return Decimal(str(value))
  28. except Exception:
  29. errors.append(f"第 {idx} 项 {field} 无效: {value!r}")
  30. return None
  31. def _optional_int_list(value: Any, field: str, idx: int, errors: list[str]) -> list[int]:
  32. if value is None:
  33. return []
  34. if not isinstance(value, list):
  35. errors.append(f"第 {idx} 项 {field} 必须是数组: {value!r}")
  36. return []
  37. out: list[int] = []
  38. for v in value:
  39. try:
  40. out.append(int(v))
  41. except (TypeError, ValueError):
  42. errors.append(f"第 {idx} 项 {field} 含无效元素: {v!r}")
  43. return []
  44. return out
  45. def _normalize_items(
  46. items: list[dict[str, Any]],
  47. *,
  48. default_biz_dt: str | None,
  49. ) -> tuple[list[dict[str, Any]], list[list[int]], list[list[int]], list[str]]:
  50. """校验并规范化待落库行,返回 (rows, related_pool_id_lists, category_id_lists, errors)。"""
  51. rows: list[dict[str, Any]] = []
  52. related_pool_id_lists: list[list[int]] = []
  53. category_id_lists: list[list[int]] = []
  54. errors: list[str] = []
  55. seen_keys: set[tuple[str, str]] = set()
  56. for idx, item in enumerate(items):
  57. if not isinstance(item, dict):
  58. errors.append(f"第 {idx} 项不是对象")
  59. continue
  60. demand_name = str(item.get("demand_name") or "").strip()
  61. if not demand_name:
  62. errors.append(f"第 {idx} 项缺少 demand_name")
  63. continue
  64. grade = str(item.get("grade") or "").strip().upper()
  65. if grade not in VALID_GRADES:
  66. errors.append(f"第 {idx} 项 grade 无效(只能是 {'/'.join(VALID_GRADES)}): {item.get('grade')!r}")
  67. continue
  68. reason = str(item.get("reason") or "").strip()
  69. if not reason:
  70. errors.append(f"第 {idx} 项缺少 reason")
  71. continue
  72. item_biz_dt, err = normalize_biz_dt(item.get("biz_dt"))
  73. if err:
  74. errors.append(f"第 {idx} 项 {err}")
  75. continue
  76. resolved_biz_dt = item_biz_dt or default_biz_dt
  77. if not resolved_biz_dt:
  78. errors.append(f"第 {idx} 项缺少 biz_dt(且未传入默认 biz_dt)")
  79. continue
  80. dedupe_key = (resolved_biz_dt, demand_name)
  81. if dedupe_key in seen_keys:
  82. errors.append(f"第 {idx} 项在本次请求中重复: biz_dt={resolved_biz_dt}, demand_name={demand_name}")
  83. continue
  84. related_pool_ids = _optional_int_list(item.get("related_pool_ids"), "related_pool_ids", idx, errors)
  85. if not related_pool_ids:
  86. errors.append(
  87. f"第 {idx} 项缺少 related_pool_ids(必填,需先用 search_related_pool_demands 找到对应的 "
  88. f"multi_demand_pool_di.id)"
  89. )
  90. continue
  91. seen_keys.add(dedupe_key)
  92. score = _optional_decimal(item.get("score"), "score", idx, errors)
  93. prior_raw = item.get("prior_total_score")
  94. if prior_raw is None or prior_raw == "" or prior_raw == "—":
  95. prior_total_score = None
  96. else:
  97. parsed_prior = _optional_decimal(prior_raw, "prior_total_score", idx, errors)
  98. prior_total_score = None if parsed_prior is None or parsed_prior == 0 else parsed_prior
  99. posterior_rov_avg = _optional_decimal(item.get("posterior_rov_avg"), "posterior_rov_avg", idx, errors)
  100. posterior_rov_count = 0
  101. if item.get("posterior_rov_count") is not None:
  102. try:
  103. posterior_rov_count = int(item.get("posterior_rov_count"))
  104. except (TypeError, ValueError):
  105. errors.append(f"第 {idx} 项 posterior_rov_count 无效: {item.get('posterior_rov_count')!r}")
  106. continue
  107. category_ids = _optional_int_list(item.get("category_ids"), "category_ids", idx, errors)
  108. has_posterior = 1 if posterior_rov_count > 0 else 0
  109. rows.append(
  110. {
  111. "biz_dt": resolved_biz_dt,
  112. "demand_name": demand_name,
  113. "category_ids": dump_int_list(category_ids),
  114. "grade": grade,
  115. "score": score,
  116. "prior_total_score": prior_total_score,
  117. "posterior_rov_avg": posterior_rov_avg,
  118. "posterior_rov_count": posterior_rov_count,
  119. "has_posterior": has_posterior,
  120. "related_pool_ids": dump_int_list(related_pool_ids),
  121. "reason": reason,
  122. }
  123. )
  124. related_pool_id_lists.append(related_pool_ids)
  125. category_id_lists.append(category_ids)
  126. return rows, related_pool_id_lists, category_id_lists, errors
  127. @tool
  128. def batch_save_demand_grades(items: list[dict[str, Any]], biz_dt: Optional[str] = None) -> str:
  129. """
  130. 批量保存需求分级结果到 demand_grade 表(按 biz_dt+demand_name upsert,可重复调用覆盖修正)。
  131. video_list(关联视频列表)与 strategies(来源策略列表)会自动从 related_pool_ids 对应的
  132. multi_demand_pool_di 原始行推导写入,无需手工传入。category_ids 除了写入展示快照字段,
  133. 也会同步写入 demand_grade_category_rel 映射表,供前端按分类高效查询。
  134. Args:
  135. items: 待保存列表,每项字段:
  136. - demand_name (必填): 需求名称
  137. - grade (必填): S/A/B/C/D 之一
  138. - reason (必填): 判断依据,需引用具体的先验/后验数值
  139. - related_pool_ids (必填): 该需求对应的 multi_demand_pool_di.id 列表,需先调用
  140. search_related_pool_demands 找到;用于关联原始需求,并自动推导 video_list/strategies
  141. - score (可选): 0-100 数值分,辅助同级排序
  142. - category_ids (可选): 归属的树节点 id 列表,会写入 demand_grade_category_rel 映射表
  143. - prior_total_score (可选): 落库时的先验 total_score 快照
  144. - posterior_rov_avg / posterior_rov_count (可选): 落库时的后验 real_rov_7d 快照;
  145. count>0 时自动标记为「有后验数据」
  146. - biz_dt (可选): 覆盖本项使用的业务日,不传则用调用时的 biz_dt 参数
  147. biz_dt: 本次调用的默认业务日期 YYYYMMDD;items 内每项也可单独指定 biz_dt 覆盖。
  148. Returns:
  149. 保存结果摘要,包含成功条数与校验失败说明。
  150. """
  151. if not items:
  152. return "items 不能为空"
  153. default_biz_dt, err = normalize_biz_dt(biz_dt)
  154. if err:
  155. return err
  156. rows, related_pool_id_lists, category_id_lists, errors = _normalize_items(
  157. items, default_biz_dt=default_biz_dt
  158. )
  159. if not rows:
  160. detail = ";".join(errors) if errors else "无有效数据"
  161. return f"没有可保存的数据: {detail}"
  162. try:
  163. with get_session() as session:
  164. pool_repo = MultiDemandPoolDiRepository(session)
  165. all_pool_ids = sorted({pid for ids in related_pool_id_lists for pid in ids})
  166. pool_rows = pool_repo.get_by_ids(all_pool_ids) if all_pool_ids else []
  167. pool_by_id = {int(r.id): r for r in pool_rows}
  168. final_rows: list[dict[str, Any]] = []
  169. saved_indices: list[int] = []
  170. for i, (row, pool_ids) in enumerate(zip(rows, related_pool_id_lists)):
  171. matched = [pool_by_id[pid] for pid in pool_ids if pid in pool_by_id]
  172. missing = [pid for pid in pool_ids if pid not in pool_by_id]
  173. if not matched:
  174. errors.append(
  175. f"demand_name={row['demand_name']!r} 的 related_pool_ids={pool_ids} "
  176. f"均未在 multi_demand_pool_di 中找到,跳过该项"
  177. )
  178. continue
  179. if missing:
  180. errors.append(
  181. f"demand_name={row['demand_name']!r} 的 related_pool_ids 中 {missing} 未找到,已忽略"
  182. )
  183. row["video_list"] = merge_video_ids(matched)
  184. row["strategies"] = collect_strategies(matched)
  185. final_rows.append(row)
  186. saved_indices.append(i)
  187. if not final_rows:
  188. detail = ";".join(errors) if errors else "无有效数据"
  189. return f"没有可保存的数据: {detail}"
  190. grade_repo = DemandGradeRepository(session)
  191. affected = grade_repo.bulk_upsert(final_rows)
  192. names_by_biz_dt: dict[str, list[str]] = {}
  193. for i in saved_indices:
  194. names_by_biz_dt.setdefault(rows[i]["biz_dt"], []).append(rows[i]["demand_name"])
  195. id_by_biz_dt_name: dict[tuple[str, str], int] = {}
  196. for bd, names in names_by_biz_dt.items():
  197. for name, demand_grade_id in grade_repo.get_ids_by_names(bd, names).items():
  198. id_by_biz_dt_name[(bd, name)] = demand_grade_id
  199. rel_repo = DemandGradeCategoryRelRepository(session)
  200. for i in saved_indices:
  201. key = (rows[i]["biz_dt"], rows[i]["demand_name"])
  202. demand_grade_id = id_by_biz_dt_name.get(key)
  203. if demand_grade_id is not None:
  204. rel_repo.replace_for_demand_grade(demand_grade_id, category_id_lists[i])
  205. parts = [f"提交 {len(rows)} 条,成功写入/更新 {affected} 条({len(final_rows)} 条通过校验)"]
  206. if errors:
  207. parts.append(f"校验失败/警告 {len(errors)} 条: " + ";".join(errors))
  208. message = "。".join(parts)
  209. logger.info("batch_save_demand_grades completed: %s", message)
  210. return message
  211. except Exception as e:
  212. logger.error("batch_save_demand_grades failed: %s", e, exc_info=True)
  213. return f"批量保存需求分级失败: {e}"