batch_save_demand_grades.py 13 KB

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