filter.py 4.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137
  1. """Formal query filter used by generated query batches."""
  2. from __future__ import annotations
  3. import hashlib
  4. import json
  5. import os
  6. from pathlib import Path
  7. import httpx
  8. from core.config import Settings, load_env_file
  9. PROMPT = Path(__file__).resolve().parent.parent / "query_filter.txt"
  10. VALID_MIN = 6
  11. def prompt_version() -> str:
  12. h = hashlib.sha256()
  13. h.update(PROMPT.read_bytes())
  14. return h.hexdigest()[:16]
  15. def filter_queries(
  16. queries: list[str],
  17. settings: Settings,
  18. *,
  19. batch: int = 40,
  20. ) -> list[dict]:
  21. """Batch-filter query strings and keep output aligned with input order."""
  22. out: list[dict] = []
  23. for i in range(0, len(queries), batch):
  24. out.extend(_filter_batch(queries[i:i + batch], settings))
  25. return out
  26. def _filter_batch(queries: list[str], settings: Settings) -> list[dict]:
  27. if not queries:
  28. return []
  29. user = json.dumps(
  30. [{"idx": i, "query": q} for i, q in enumerate(queries)],
  31. ensure_ascii=False,
  32. )
  33. messages = [
  34. {"role": "system", "content": PROMPT.read_text("utf-8")},
  35. {"role": "user", "content": user},
  36. ]
  37. try:
  38. txt = _chat_content(settings, messages)
  39. data = json.loads(txt)
  40. arr = data if isinstance(data, list) else next(
  41. (v for v in data.values() if isinstance(v, list)),
  42. [],
  43. )
  44. by = {d.get("idx"): d for d in arr if isinstance(d, dict)}
  45. out = []
  46. for i in range(len(queries)):
  47. d = by.get(i, {})
  48. try:
  49. valid = int(d.get("valid", 10))
  50. except (TypeError, ValueError):
  51. valid = 10
  52. relevant = bool(d.get("relevant", True))
  53. out.append(
  54. {
  55. "keep": valid >= VALID_MIN and relevant,
  56. "valid": valid,
  57. "relevant": relevant,
  58. "reason": str(d.get("reason", ""))[:50],
  59. }
  60. )
  61. return out
  62. except Exception as exc:
  63. reason = f"筛选失败:{str(exc)[:30]}"
  64. return [
  65. {"keep": True, "valid": 10, "relevant": True, "reason": reason}
  66. for _ in queries
  67. ]
  68. def _chat_content(settings: Settings, messages: list[dict]) -> str:
  69. body = {
  70. "model": settings.llm_model,
  71. "messages": messages,
  72. "response_format": {"type": "json_object"},
  73. }
  74. env = load_env_file(os.getenv("CK_ENV_FILE", ".env"))
  75. prefer_ark = os.getenv("QUERY_FILTER_PROVIDER") == "ark" or bool(
  76. os.getenv("ARK_CHAT_MODEL")
  77. )
  78. openrouter_exc: Exception | None = None
  79. if settings.openrouter_api_key and not prefer_ark:
  80. try:
  81. resp = httpx.post(
  82. settings.openrouter_base_url.rstrip("/") + "/chat/completions",
  83. headers={
  84. "Authorization": f"Bearer {settings.openrouter_api_key}",
  85. "Content-Type": "application/json",
  86. },
  87. json=body,
  88. timeout=120,
  89. )
  90. resp.raise_for_status()
  91. return resp.json()["choices"][0]["message"]["content"]
  92. except Exception as exc:
  93. openrouter_exc = exc
  94. ark_key = os.getenv("ARK_API_KEY") or env.get("ARK_API_KEY")
  95. if ark_key:
  96. ark_model = (
  97. os.getenv("ARK_CHAT_MODEL")
  98. or env.get("ARK_CHAT_MODEL")
  99. or "doubao-seed-1-6-flash-250615"
  100. )
  101. ark_url = (
  102. os.getenv("ARK_CHAT_URL")
  103. or env.get("ARK_CHAT_URL")
  104. or "https://ark.cn-beijing.volces.com/api/v3/chat/completions"
  105. )
  106. resp = httpx.post(
  107. ark_url,
  108. headers={
  109. "Authorization": f"Bearer {ark_key}",
  110. "Content-Type": "application/json",
  111. },
  112. json={
  113. "model": ark_model,
  114. "messages": messages,
  115. "response_format": {"type": "json_object"},
  116. },
  117. timeout=120,
  118. )
  119. resp.raise_for_status()
  120. return resp.json()["choices"][0]["message"]["content"]
  121. if openrouter_exc:
  122. raise openrouter_exc
  123. raise RuntimeError("missing OpenRouter/Ark chat credentials")