filter.py 4.3 KB

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