Files
pod_trend_agent/graph/nodes/pinterest_search_node.py
T

104 lines
4.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Pinterest 参考模式节点 1/3LLM 生成搜索词(pinterest_search)。
流程:国家 Pinterest 种子词池 → LLM 生成搜索词(json_schema 结构化 + 动态注入已用词防重复)
→ 全局过滤(已用/黑名单/不适合T恤/去重)→ 持久化已用词。
兜底链:LLM json_schema → json_object → 解析失败/调用失败 → 回退种子词池随机抽样。
带 with_fallback:任何异常都不中断,返回空列表由下游跳过。
"""
import random
from typing import Any, Dict, List
from graph.llms import get_backend
from graph.pinterest import (
filter_search_terms,
load_used_terms,
merge_used,
sample_seeds,
save_used_terms,
)
from graph.validate import with_fallback
@with_fallback("pinterest_search")
def pinterest_search_node(state: Dict[str, Any]) -> Dict[str, Any]:
country = state["country"]
config = state["config"]
output_dir = state["output_dir"]
errors = list(state.get("errors") or [])
pcfg = config.get("pinterest") or {}
if not pcfg.get("enabled", True):
return {"pinterest_search_terms": [], "errors": errors}
provider = str(pcfg.get("provider") or "openai").strip().lower()
want = int(pcfg.get("search_terms_per_run", 10))
seed_sample = int(pcfg.get("seed_sample", 40))
max_used_in_prompt = int(pcfg.get("max_used_terms_in_prompt", 100))
blacklist = config.get("blacklist") or []
# 1) 种子词池(随机抽样)+ 已用搜索词
seeds = sample_seeds(country, seed_sample)
used = load_used_terms(output_dir, country)
if not seeds:
print(f"[pinterest_search] {country} 无种子词,跳过搜索词生成")
return {"pinterest_search_terms": [], "errors": errors}
# 2) LLM 生成(json_schema + 动态注入已用词)
# 已用词只取最近 N 个(默认 100)注入提示词,防 token 超限;过滤仍用全量。
used_llm = used[-max_used_in_prompt:] if max_used_in_prompt > 0 else []
terms: List[str] = []
llm = None
if provider != "static":
try:
llm = get_backend(provider)
if hasattr(llm, "bind_config"):
llm.bind_config(config.get("llm_screen") or {})
if provider not in ("mock",) and not getattr(llm, "has_key", False):
print(f"[pinterest_search] {provider} 未配置 API key,降级 mock")
llm = get_backend("mock")
except Exception as e: # noqa: BLE001
print(f"[pinterest_search] LLM 初始化失败: {e}")
llm = None
if llm is not None and hasattr(llm, "generate_pinterest_terms"):
try:
ctx = {"country": country, "seeds": seeds, "used_terms": used_llm, "count": want}
res = llm.generate_pinterest_terms(ctx)
terms = [str(t).strip() for t in (res.get("search_terms") or []) if str(t).strip()]
print(f"[pinterest_search] LLM 生成搜索词 {len(terms)} 个({country},已用词注入 {len(used_llm)}/{len(used)}")
except Exception as e: # noqa: BLE001
print(f"[pinterest_search] LLM 生成失败,回退种子词池: {e}")
terms = []
# 3) 兜底:LLM 无结果 → 种子词池随机抽样
if not terms:
terms = random.sample(seeds, min(want, len(seeds))) if seeds else []
print(f"[pinterest_search] 兜底:从种子词池取 {len(terms)} 个")
# 4) 全局过滤(已用/黑名单/不适合T恤/去重)
filtered = filter_search_terms(terms, used, blacklist)
if len(filtered) < want and seeds:
# 不足时用种子词池补充(同样过滤),保证数量
extra = filter_search_terms(seeds, merge_used(used, filtered), blacklist)
for t in extra:
if len(filtered) >= want:
break
filtered.append(t)
# 5) 持久化已用词
new_used = merge_used(used, filtered)
save_used_terms(output_dir, country, new_used)
stats = dict(state.get("stats") or {})
stats["pinterest_search"] = {
"provider": provider,
"generated": len(terms),
"filtered": len(filtered),
"used_total": len(new_used),
}
print(f"[pinterest_search] 搜索词 {len(filtered)} 个(已用累计 {len(new_used)}: "
f"{', '.join(filtered[:6])}{'...' if len(filtered) > 6 else ''}")
return {"pinterest_search_terms": filtered, "stats": stats, "errors": errors}