""" 在需求池中按同名/包含关系搜索需求词,用于合并同语义、措辞不同的需求一起判断。 """ 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()