"""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 from core.text_limits import ( ERROR_MESSAGE_MAX_CHARS, QUERY_FILTER_REASON_MAX_CHARS, clip_text, ) 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": clip_text(d.get("reason", ""), QUERY_FILTER_REASON_MAX_CHARS), } ) return out except Exception as exc: reason = f"筛选失败:{clip_text(exc, ERROR_MESSAGE_MAX_CHARS)}" 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")