| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137 |
- """Formal query filter used by generated query batches."""
- from __future__ import annotations
- import hashlib
- import json
- import os
- from pathlib import Path
- import httpx
- from core.config import Settings, load_env_file
- PROMPT = Path(__file__).resolve().parent.parent / "query_filter.txt"
- VALID_MIN = 6
- def prompt_version() -> str:
- h = hashlib.sha256()
- h.update(PROMPT.read_bytes())
- return h.hexdigest()[:16]
- def filter_queries(
- queries: list[str],
- settings: Settings,
- *,
- batch: int = 40,
- ) -> list[dict]:
- """Batch-filter query strings and keep output aligned with input order."""
- out: list[dict] = []
- for i in range(0, len(queries), batch):
- out.extend(_filter_batch(queries[i:i + batch], settings))
- return out
- def _filter_batch(queries: list[str], settings: Settings) -> list[dict]:
- if not queries:
- return []
- user = json.dumps(
- [{"idx": i, "query": q} for i, q in enumerate(queries)],
- ensure_ascii=False,
- )
- messages = [
- {"role": "system", "content": PROMPT.read_text("utf-8")},
- {"role": "user", "content": user},
- ]
- try:
- txt = _chat_content(settings, messages)
- data = json.loads(txt)
- arr = data if isinstance(data, list) else next(
- (v for v in data.values() if isinstance(v, list)),
- [],
- )
- by = {d.get("idx"): d for d in arr if isinstance(d, dict)}
- out = []
- for i in range(len(queries)):
- d = by.get(i, {})
- try:
- valid = int(d.get("valid", 10))
- except (TypeError, ValueError):
- valid = 10
- relevant = bool(d.get("relevant", True))
- out.append(
- {
- "keep": valid >= VALID_MIN and relevant,
- "valid": valid,
- "relevant": relevant,
- "reason": str(d.get("reason", ""))[:50],
- }
- )
- return out
- except Exception as exc:
- reason = f"筛选失败:{str(exc)[:30]}"
- return [
- {"keep": True, "valid": 10, "relevant": True, "reason": reason}
- for _ in queries
- ]
- def _chat_content(settings: Settings, messages: list[dict]) -> str:
- body = {
- "model": settings.llm_model,
- "messages": messages,
- "response_format": {"type": "json_object"},
- }
- env = load_env_file(os.getenv("CK_ENV_FILE", ".env"))
- prefer_ark = os.getenv("QUERY_FILTER_PROVIDER") == "ark" or bool(
- os.getenv("ARK_CHAT_MODEL")
- )
- openrouter_exc: Exception | None = None
- if settings.openrouter_api_key and not prefer_ark:
- try:
- resp = httpx.post(
- settings.openrouter_base_url.rstrip("/") + "/chat/completions",
- headers={
- "Authorization": f"Bearer {settings.openrouter_api_key}",
- "Content-Type": "application/json",
- },
- json=body,
- timeout=120,
- )
- resp.raise_for_status()
- return resp.json()["choices"][0]["message"]["content"]
- except Exception as exc:
- openrouter_exc = exc
- ark_key = os.getenv("ARK_API_KEY") or env.get("ARK_API_KEY")
- if ark_key:
- ark_model = (
- os.getenv("ARK_CHAT_MODEL")
- or env.get("ARK_CHAT_MODEL")
- or "doubao-seed-1-6-flash-250615"
- )
- ark_url = (
- os.getenv("ARK_CHAT_URL")
- or env.get("ARK_CHAT_URL")
- or "https://ark.cn-beijing.volces.com/api/v3/chat/completions"
- )
- resp = httpx.post(
- ark_url,
- headers={
- "Authorization": f"Bearer {ark_key}",
- "Content-Type": "application/json",
- },
- json={
- "model": ark_model,
- "messages": messages,
- "response_format": {"type": "json_object"},
- },
- timeout=120,
- )
- resp.raise_for_status()
- return resp.json()["choices"][0]["message"]["content"]
- if openrouter_exc:
- raise openrouter_exc
- raise RuntimeError("missing OpenRouter/Ark chat credentials")
|