| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117 |
- """
- 在需求池中按同名/包含关系搜索需求词,用于合并同语义、措辞不同的需求一起判断。
- """
- from __future__ import annotations
- import logging
- from sqlalchemy.orm import Session
- from agents.demand_grade_agent.tools.demand_priority import (
- build_demand_priority_index,
- format_demand_priority,
- )
- from agents.demand_grade_agent.tools.shared import (
- format_posterior_value,
- format_score,
- normalize_biz_dt,
- normalize_str_list,
- )
- from supply_agent.tools import tool
- from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
- from supply_infra.db.session import get_session
- logger = logging.getLogger(__name__)
- def _search_one_related_pool_demands(
- session: Session,
- normalized: str,
- keyword: str,
- priority_index: dict[str, dict],
- ) -> str:
- rows = MultiDemandPoolDiRepository(session).search_rows_by_name_fragment(normalized, keyword)
- if not rows:
- return f"biz_dt={normalized} 未找到与「{keyword}」同名/包含关系的需求词"
- lines = []
- for demand_name in dict.fromkeys(str(row["demand_name"]) for row in rows):
- lines.append(f"需求自身证据「{demand_name}」:")
- lines.extend(format_demand_priority(priority_index.get(demand_name)))
- for row in rows:
- rov = format_posterior_value(row["real_rov_7d"])
- vov = format_posterior_value(row["real_vov_7d"])
- weight = format_score(row["weight"])
- video_count = row["video_count"] if row["video_count"] is not None else "—"
- lines.append(
- f"[id={row['id']}|{row['strategy']}|weight={weight}"
- f"|视频数={video_count}|rov_diff={rov}|vov_diff={vov}] {row['demand_name']}"
- )
- return "\n".join(lines)
- @tool
- def search_related_pool_demands(biz_dt: str, keywords: list[str]) -> str:
- """
- 按同名/包含关系搜索需求池,合并同语义需求的多条记录一起判断。
- 支持批量传入多个 keyword,同一 biz_dt 下一次调用返回各关键词的匹配结果;
- 每段结果前会标注原始 keyword。
- 匹配规则为双向包含:keyword 是 demand_name 的子串,或 demand_name 是 keyword 的子串。
- 用于发现措辞不同但语义相同/高度相关的需求词(例如「减脂期加餐」与「减脂加餐」),
- 在判级时应把这些记录一起纳入参考,而不是只看单条记录。
- Args:
- biz_dt: 业务日期,格式 YYYYMMDD(必填)。
- keywords: 需求词或关键片段列表,可一次传多个。
- Returns:
- 每个 keyword 一段,段首标注 `--- keyword: xxx ---`,例如:
- --- keyword: 加餐 ---
- [id=101|strategy_a|weight=3.20|视频数=12|rov_diff=0.0500 [效果非常好]|vov_diff=-0.0800 [可接受]] 减脂期加餐怎么吃
- """
- normalized, err = normalize_biz_dt(biz_dt)
- if err:
- return err
- if not normalized:
- return "biz_dt 不能为空"
- keyword_list, err = normalize_str_list(keywords, "keywords")
- if err:
- return err
- try:
- with get_session() as session:
- pool_repo = MultiDemandPoolDiRepository(session)
- priority_index = build_demand_priority_index(pool_repo.list_by_biz_dt(normalized))
- sections: list[str] = []
- for keyword in keyword_list:
- result = _search_one_related_pool_demands(
- session,
- normalized,
- keyword,
- priority_index,
- )
- sections.append(f"--- keyword: {keyword} ---\n{result}")
- message = "\n\n".join(sections)
- logger.info(
- "search_related_pool_demands completed: biz_dt=%s count=%d",
- normalized,
- len(keyword_list),
- )
- return message
- except Exception as e:
- logger.error("search_related_pool_demands failed: %s", e, exc_info=True)
- return f"搜索关联需求词失败: {e}"
- def main() -> None:
- print(search_related_pool_demands(biz_dt="20260714", keywords=["加餐", "减脂"]))
- if __name__ == "__main__":
- main()
|