search_related_pool_demands.py 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117
  1. """
  2. 在需求池中按同名/包含关系搜索需求词,用于合并同语义、措辞不同的需求一起判断。
  3. """
  4. from __future__ import annotations
  5. import logging
  6. from sqlalchemy.orm import Session
  7. from agents.demand_grade_agent.tools.demand_priority import (
  8. build_demand_priority_index,
  9. format_demand_priority,
  10. )
  11. from agents.demand_grade_agent.tools.shared import (
  12. format_posterior_value,
  13. format_score,
  14. normalize_biz_dt,
  15. normalize_str_list,
  16. )
  17. from supply_agent.tools import tool
  18. from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
  19. from supply_infra.db.session import get_session
  20. logger = logging.getLogger(__name__)
  21. def _search_one_related_pool_demands(
  22. session: Session,
  23. normalized: str,
  24. keyword: str,
  25. priority_index: dict[str, dict],
  26. ) -> str:
  27. rows = MultiDemandPoolDiRepository(session).search_rows_by_name_fragment(normalized, keyword)
  28. if not rows:
  29. return f"biz_dt={normalized} 未找到与「{keyword}」同名/包含关系的需求词"
  30. lines = []
  31. for demand_name in dict.fromkeys(str(row["demand_name"]) for row in rows):
  32. lines.append(f"需求自身证据「{demand_name}」:")
  33. lines.extend(format_demand_priority(priority_index.get(demand_name)))
  34. for row in rows:
  35. rov = format_posterior_value(row["real_rov_7d"])
  36. vov = format_posterior_value(row["real_vov_7d"])
  37. weight = format_score(row["weight"])
  38. video_count = row["video_count"] if row["video_count"] is not None else "—"
  39. lines.append(
  40. f"[id={row['id']}|{row['strategy']}|weight={weight}"
  41. f"|视频数={video_count}|rov_diff={rov}|vov_diff={vov}] {row['demand_name']}"
  42. )
  43. return "\n".join(lines)
  44. @tool
  45. def search_related_pool_demands(biz_dt: str, keywords: list[str]) -> str:
  46. """
  47. 按同名/包含关系搜索需求池,合并同语义需求的多条记录一起判断。
  48. 支持批量传入多个 keyword,同一 biz_dt 下一次调用返回各关键词的匹配结果;
  49. 每段结果前会标注原始 keyword。
  50. 匹配规则为双向包含:keyword 是 demand_name 的子串,或 demand_name 是 keyword 的子串。
  51. 用于发现措辞不同但语义相同/高度相关的需求词(例如「减脂期加餐」与「减脂加餐」),
  52. 在判级时应把这些记录一起纳入参考,而不是只看单条记录。
  53. Args:
  54. biz_dt: 业务日期,格式 YYYYMMDD(必填)。
  55. keywords: 需求词或关键片段列表,可一次传多个。
  56. Returns:
  57. 每个 keyword 一段,段首标注 `--- keyword: xxx ---`,例如:
  58. --- keyword: 加餐 ---
  59. [id=101|strategy_a|weight=3.20|视频数=12|rov_diff=0.0500 [效果非常好]|vov_diff=-0.0800 [可接受]] 减脂期加餐怎么吃
  60. """
  61. normalized, err = normalize_biz_dt(biz_dt)
  62. if err:
  63. return err
  64. if not normalized:
  65. return "biz_dt 不能为空"
  66. keyword_list, err = normalize_str_list(keywords, "keywords")
  67. if err:
  68. return err
  69. try:
  70. with get_session() as session:
  71. pool_repo = MultiDemandPoolDiRepository(session)
  72. priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(normalized))
  73. sections: list[str] = []
  74. for keyword in keyword_list:
  75. result = _search_one_related_pool_demands(
  76. session,
  77. normalized,
  78. keyword,
  79. priority_index,
  80. )
  81. sections.append(f"--- keyword: {keyword} ---\n{result}")
  82. message = "\n\n".join(sections)
  83. logger.info(
  84. "search_related_pool_demands completed: biz_dt=%s count=%d",
  85. normalized,
  86. len(keyword_list),
  87. )
  88. return message
  89. except Exception as e:
  90. logger.error("search_related_pool_demands failed: %s", e, exc_info=True)
  91. return f"搜索关联需求词失败: {e}"
  92. def main() -> None:
  93. print(search_related_pool_demands(biz_dt="20260714", keywords=["加餐", "减脂"]))
  94. if __name__ == "__main__":
  95. main()