batch_save_generated_demands.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286
  1. """
  2. 批量保存单维度需求产生结果到 generated_demand 表。
  3. """
  4. from __future__ import annotations
  5. import logging
  6. import uuid
  7. from decimal import Decimal
  8. from typing import Any
  9. from agents.generate_demand_agent.tools.dim_constants import (
  10. DIM_KEYS,
  11. DIM_LABEL,
  12. build_category_path,
  13. resolve_biz_dt,
  14. )
  15. from supply_agent.tools import tool
  16. from supply_infra.db.repositories.demand_belong_category_repo import (
  17. DemandBelongCategoryRepository,
  18. )
  19. from supply_infra.db.repositories.demand_popularity_stats_repo import (
  20. DemandPopularityStatsRepository,
  21. )
  22. from supply_infra.db.repositories.generated_demand_repo import GeneratedDemandRepository
  23. from supply_infra.db.repositories.global_tree_category_repo import (
  24. GlobalTreeCategoryRepository,
  25. )
  26. from supply_infra.db.session import get_session
  27. logger = logging.getLogger(__name__)
  28. def _optional_int(value: Any) -> int | None:
  29. if value is None or value == "":
  30. return None
  31. return int(value)
  32. def _optional_str(value: Any) -> str | None:
  33. if value is None:
  34. return None
  35. text = str(value).strip()
  36. return text or None
  37. def _normalize_items(
  38. items: list[dict[str, Any]],
  39. *,
  40. default_run_id: str,
  41. ) -> tuple[list[dict[str, Any]], list[str]]:
  42. """校验并规范化待插入行,返回 (rows, errors)。"""
  43. rows: list[dict[str, Any]] = []
  44. errors: list[str] = []
  45. seen_keys: set[tuple[str, str]] = set()
  46. for idx, item in enumerate(items):
  47. if not isinstance(item, dict):
  48. errors.append(f"第 {idx} 项不是对象")
  49. continue
  50. source_dim = _optional_str(item.get("source_dim"))
  51. overall_direction = _optional_str(item.get("overall_direction"))
  52. summary_event = _optional_str(item.get("summary_event"))
  53. demand_name = _optional_str(item.get("demand_name"))
  54. if not source_dim:
  55. errors.append(f"第 {idx} 项缺少 source_dim")
  56. continue
  57. if source_dim not in DIM_KEYS:
  58. allowed = "、".join(f"{k}({DIM_LABEL[k]})" for k in DIM_KEYS)
  59. errors.append(f"第 {idx} 项 source_dim 无效(只能是:{allowed}): {source_dim}")
  60. continue
  61. if not overall_direction:
  62. errors.append(f"第 {idx} 项缺少 overall_direction")
  63. continue
  64. if not summary_event:
  65. errors.append(f"第 {idx} 项缺少 summary_event")
  66. continue
  67. if not demand_name:
  68. errors.append(f"第 {idx} 项缺少 demand_name")
  69. continue
  70. dedupe_key = (source_dim, demand_name)
  71. if dedupe_key in seen_keys:
  72. errors.append(
  73. f"第 {idx} 项在本次请求中重复: source_dim={source_dim}, demand_name={demand_name}"
  74. )
  75. continue
  76. seen_keys.add(dedupe_key)
  77. try:
  78. category_id = _optional_int(item.get("category_id"))
  79. except (TypeError, ValueError):
  80. errors.append(f"第 {idx} 项 category_id 无效: {item.get('category_id')!r}")
  81. continue
  82. try:
  83. demand_belong_id = _optional_int(item.get("demand_belong_id"))
  84. except (TypeError, ValueError):
  85. errors.append(
  86. f"第 {idx} 项 demand_belong_id 无效: {item.get('demand_belong_id')!r}"
  87. )
  88. continue
  89. run_id = _optional_str(item.get("run_id")) or default_run_id
  90. reason = _optional_str(item.get("reason"))
  91. category_path = _optional_str(item.get("category_path"))
  92. biz_dt = _optional_str(item.get("biz_dt"))
  93. dim_avg: Decimal | None = None
  94. dim_count: int | None = None
  95. if item.get("dim_avg") is not None and item.get("dim_avg") != "":
  96. try:
  97. dim_avg = Decimal(str(item.get("dim_avg")))
  98. except Exception:
  99. errors.append(f"第 {idx} 项 dim_avg 无效: {item.get('dim_avg')!r}")
  100. continue
  101. if item.get("dim_count") is not None and item.get("dim_count") != "":
  102. try:
  103. dim_count = int(item.get("dim_count"))
  104. except (TypeError, ValueError):
  105. errors.append(f"第 {idx} 项 dim_count 无效: {item.get('dim_count')!r}")
  106. continue
  107. rows.append(
  108. {
  109. "source_dim": source_dim,
  110. "overall_direction": overall_direction,
  111. "summary_event": summary_event,
  112. "demand_name": demand_name,
  113. "demand_belong_id": demand_belong_id,
  114. "category_id": category_id,
  115. "category_path": category_path,
  116. "dim_avg": dim_avg,
  117. "dim_count": dim_count,
  118. "reason": reason,
  119. "biz_dt": biz_dt,
  120. "run_id": run_id,
  121. "is_delete": 0,
  122. }
  123. )
  124. return rows, errors
  125. @tool
  126. def batch_save_generated_demands(
  127. items: list[dict[str, Any]],
  128. biz_dt: str | None = None,
  129. ) -> str:
  130. """
  131. 批量保存单维度需求产生结果。
  132. 写入 generated_demand 表。demand_name 必须已存在于 demand_belong_category;
  133. 不存在的词会被跳过。同一 run_id 内 (source_dim, demand_name) 去重。
  134. Args:
  135. items: 待保存列表。每项必填:
  136. - source_dim: 四维之一(ext_pop / plat_sust_pop / plat_ly_pop / recent_pop)
  137. - overall_direction: 整体方向
  138. - summary_event: 汇总事件
  139. - demand_name: 需求名(须来自 demand_belong_category.name)
  140. 选填:
  141. - category_id, category_path, demand_belong_id, reason,
  142. dim_avg, dim_count, biz_dt, run_id
  143. 若省略 run_id,本次调用自动生成同一 run_id。
  144. 若省略 demand_belong_id / category_id / 热度快照,将尽量从库中补全。
  145. biz_dt: 业务日 YYYYMMDD;省略则使用 demand_popularity_stats 最新业务日。
  146. 会作为本次落库的默认 biz_dt,并用于补全热度快照。
  147. Returns:
  148. 保存结果摘要。
  149. """
  150. if not items:
  151. return "items 不能为空"
  152. default_run_id = uuid.uuid4().hex
  153. rows, errors = _normalize_items(items, default_run_id=default_run_id)
  154. if not rows:
  155. detail = ";".join(errors) if errors else "无有效数据"
  156. return f"没有可保存的数据: {detail}"
  157. try:
  158. with get_session() as session:
  159. stats_repo = DemandPopularityStatsRepository(session)
  160. resolved_biz_dt, err = resolve_biz_dt(
  161. biz_dt,
  162. get_latest=stats_repo.get_latest_biz_dt,
  163. has_data=stats_repo.has_biz_dt,
  164. table_label="demand_popularity_stats 数据",
  165. )
  166. if err:
  167. return err
  168. belong_repo = DemandBelongCategoryRepository(session)
  169. names = [r["demand_name"] for r in rows]
  170. belong_by_name = belong_repo.get_by_names(names)
  171. valid_rows: list[dict[str, Any]] = []
  172. skipped_names: list[str] = []
  173. for row in rows:
  174. belong = belong_by_name.get(row["demand_name"])
  175. if belong is None:
  176. skipped_names.append(row["demand_name"])
  177. continue
  178. if row["demand_belong_id"] is None:
  179. row["demand_belong_id"] = int(belong.id)
  180. if row["category_id"] is None:
  181. row["category_id"] = int(belong.category_id)
  182. valid_rows.append(row)
  183. if not valid_rows:
  184. detail = "、".join(skipped_names)
  185. extra = f";校验失败: {';'.join(errors)}" if errors else ""
  186. return f"没有可保存的数据: demand_name 不存在于 demand_belong_category: {detail}{extra}"
  187. # 补全路径
  188. need_path_ids = [
  189. int(r["category_id"])
  190. for r in valid_rows
  191. if r["category_id"] is not None and not r.get("category_path")
  192. ]
  193. if need_path_ids:
  194. categories = GlobalTreeCategoryRepository(session).list_active_categories()
  195. by_id = {int(c.id): c for c in categories}
  196. for row in valid_rows:
  197. if row.get("category_path") or row["category_id"] is None:
  198. continue
  199. path = build_category_path(int(row["category_id"]), by_id)
  200. if path:
  201. row["category_path"] = path
  202. # 补全 biz_dt 与热度快照
  203. need_stats_ids = [
  204. int(r["demand_belong_id"])
  205. for r in valid_rows
  206. if r["demand_belong_id"] is not None
  207. and (r.get("dim_avg") is None or r.get("dim_count") is None)
  208. ]
  209. stats_by_belong: dict[int, Any] = {}
  210. if need_stats_ids:
  211. for stats in stats_repo.list_by_biz_dt_and_belong_ids(
  212. resolved_biz_dt, need_stats_ids
  213. ):
  214. stats_by_belong[int(stats.demand_category_id)] = stats
  215. for row in valid_rows:
  216. if not row.get("biz_dt"):
  217. row["biz_dt"] = resolved_biz_dt
  218. belong_id = row.get("demand_belong_id")
  219. if belong_id is None:
  220. continue
  221. stats = stats_by_belong.get(int(belong_id))
  222. if stats is None:
  223. continue
  224. dim = row["source_dim"]
  225. if row.get("dim_count") is None:
  226. row["dim_count"] = int(getattr(stats, f"{dim}_count", 0) or 0)
  227. if row.get("dim_avg") is None:
  228. avg = getattr(stats, f"{dim}_avg", None)
  229. row["dim_avg"] = Decimal(str(avg)) if avg is not None else None
  230. inserted = GeneratedDemandRepository(session).bulk_insert(valid_rows)
  231. run_ids = sorted({str(r["run_id"]) for r in valid_rows})
  232. parts = [
  233. f"提交有效 {len(valid_rows)} 条,成功插入 {inserted} 条",
  234. f"run_id={','.join(run_ids)}",
  235. f"biz_dt={resolved_biz_dt}",
  236. ]
  237. if skipped_names:
  238. parts.append(
  239. f"因 demand_name 不存在跳过 {len(skipped_names)} 条: "
  240. + "、".join(skipped_names[:20])
  241. + ("…" if len(skipped_names) > 20 else "")
  242. )
  243. if errors:
  244. parts.append(f"校验失败 {len(errors)} 条: " + ";".join(errors[:20]))
  245. message = "。".join(parts)
  246. logger.info("batch_save_generated_demands completed: %s", message)
  247. return message
  248. except Exception as e:
  249. logger.error("batch_save_generated_demands failed: %s", e, exc_info=True)
  250. return f"批量保存生成需求失败: {e}"