query_filter.py 2.2 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546
  1. """创作 query 筛选器:调用 acquisition/query_filter.txt(LLM 只做 keep/排除),批量判一组 query。
  2. 命中 A–E 任一即排除:A 语义不通 / B 题材本身 / C 制作工具 / D 应试学术公务政企 / E 作品素材。
  3. 返回与输入等长的 [{keep: bool, reason: str}]。供 build_creation_demo、filter_multiaxis 等共用。
  4. """
  5. from __future__ import annotations
  6. import json
  7. from pathlib import Path
  8. import httpx
  9. from core.config import Settings
  10. PROMPT = Path(__file__).resolve().parent / "query_filter.txt"
  11. def filter_queries(queries: list[str], settings: Settings, *, batch: int = 40) -> list[dict]:
  12. """批量过滤(分批调用,避免一次喂太多)。返回 [{keep, reason}],与 queries 对齐。"""
  13. out: list[dict] = []
  14. for i in range(0, len(queries), batch):
  15. out.extend(_filter_batch(queries[i:i + batch], settings))
  16. return out
  17. def _filter_batch(queries: list[str], settings: Settings) -> list[dict]:
  18. if not queries:
  19. return []
  20. user = json.dumps([{"idx": i, "query": q} for i, q in enumerate(queries)], ensure_ascii=False)
  21. api = settings.openrouter_base_url.rstrip("/") + "/chat/completions"
  22. headers = {"Authorization": f"Bearer {settings.openrouter_api_key}", "Content-Type": "application/json"}
  23. body = {"model": settings.llm_model, "messages": [
  24. {"role": "system", "content": PROMPT.read_text("utf-8")},
  25. {"role": "user", "content": user}], "response_format": {"type": "json_object"}}
  26. try:
  27. resp = httpx.post(api, headers=headers, json=body, timeout=120)
  28. resp.raise_for_status()
  29. txt = resp.json()["choices"][0]["message"]["content"]
  30. # query_filter 要求输出数组;有的模型会包一层 {"result":[...]},都兜住
  31. data = json.loads(txt)
  32. arr = data if isinstance(data, list) else next((v for v in data.values() if isinstance(v, list)), [])
  33. by = {d.get("idx"): d for d in arr if isinstance(d, dict)}
  34. return [{"keep": bool(by.get(i, {}).get("keep", True)),
  35. "reason": str(by.get(i, {}).get("reason", ""))[:50]} for i in range(len(queries))]
  36. except Exception as exc:
  37. return [{"keep": True, "reason": f"筛选失败:{str(exc)[:30]}"} for _ in queries]