106 lines
3.2 KiB
Python
106 lines
3.2 KiB
Python
|
|
"""LLM-реранкер через OpenCode/DeepSeek.
|
|||
|
|
|
|||
|
|
Отправляет top-N кандидатов с промптом «верни JSON-список rowid, отсортированных
|
|||
|
|
по релевантности». При любой ошибке возвращает ``None`` → caller использует
|
|||
|
|
нереранкнутый список.
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
from __future__ import annotations
|
|||
|
|
|
|||
|
|
import json
|
|||
|
|
import logging
|
|||
|
|
import os
|
|||
|
|
import re
|
|||
|
|
from typing import List, Optional
|
|||
|
|
|
|||
|
|
from .bm25 import Hit
|
|||
|
|
|
|||
|
|
logger = logging.getLogger(__name__)
|
|||
|
|
|
|||
|
|
|
|||
|
|
RERANK_PROMPT = """Ты — реранкер для поисковой выдачи. Тебе дан запрос и {n} фрагментов документов.
|
|||
|
|
Верни JSON-список ``rowid`` (целые числа) В ПОРЯДКЕ убывания релевантности запросу.
|
|||
|
|
Не добавляй пояснений, только JSON.
|
|||
|
|
|
|||
|
|
Запрос: {query}
|
|||
|
|
|
|||
|
|
Фрагменты:
|
|||
|
|
{chunks}
|
|||
|
|
|
|||
|
|
Верни ТОЛЬКО JSON-массив rowid, например: ``[42, 17, 5]``
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
|
|||
|
|
def llm_rerank(
|
|||
|
|
query: str,
|
|||
|
|
hits: List[Hit],
|
|||
|
|
api_key: str = "",
|
|||
|
|
base_url: str = "https://opencode.ai/zen/v1",
|
|||
|
|
model: str = "deepseek-v4-flash-free",
|
|||
|
|
top_k: int = 20,
|
|||
|
|
) -> Optional[List[Hit]]:
|
|||
|
|
"""Отправляет ``top_k`` чанков в LLM и возвращает пересортированный список.
|
|||
|
|
|
|||
|
|
Возвращает ``None`` если запрос не удался.
|
|||
|
|
"""
|
|||
|
|
if not hits:
|
|||
|
|
return []
|
|||
|
|
api_key = api_key or os.environ.get("OPENCODE_API_KEY", "")
|
|||
|
|
if not api_key:
|
|||
|
|
logger.warning("[rerank] OPENCODE_API_KEY not set, skipping")
|
|||
|
|
return None
|
|||
|
|
|
|||
|
|
candidates = hits[:top_k]
|
|||
|
|
chunks_text = "\n\n".join(
|
|||
|
|
f"rowid={h.rowid}: {h.snippet(400)}" for h in candidates
|
|||
|
|
)
|
|||
|
|
prompt = RERANK_PROMPT.format(n=len(candidates), query=query, chunks=chunks_text)
|
|||
|
|
|
|||
|
|
try:
|
|||
|
|
from openai import AsyncOpenAI
|
|||
|
|
client = AsyncOpenAI(base_url=base_url, api_key=api_key)
|
|||
|
|
response = client.chat.completions.create(
|
|||
|
|
model=model,
|
|||
|
|
messages=[{"role": "user", "content": prompt}],
|
|||
|
|
temperature=0.0,
|
|||
|
|
max_tokens=512,
|
|||
|
|
)
|
|||
|
|
content = (response.choices[0].message.content or "").strip()
|
|||
|
|
order_ids = _parse_ids(content)
|
|||
|
|
if not order_ids:
|
|||
|
|
return None
|
|||
|
|
except Exception as exc:
|
|||
|
|
logger.warning("[rerank] LLM call failed: %s", exc)
|
|||
|
|
return None
|
|||
|
|
|
|||
|
|
by_id = {h.rowid: h for h in candidates}
|
|||
|
|
result: List[Hit] = []
|
|||
|
|
for rid in order_ids:
|
|||
|
|
hit = by_id.get(int(rid))
|
|||
|
|
if hit is None:
|
|||
|
|
continue
|
|||
|
|
result.append(hit)
|
|||
|
|
for h in candidates:
|
|||
|
|
if h.rowid not in {r.rowid for r in result}:
|
|||
|
|
result.append(h)
|
|||
|
|
return result
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _parse_ids(content: str) -> List[int]:
|
|||
|
|
match = re.search(r"\[[^\]]*\]", content, re.DOTALL)
|
|||
|
|
if not match:
|
|||
|
|
return []
|
|||
|
|
try:
|
|||
|
|
data = json.loads(match.group(0))
|
|||
|
|
except json.JSONDecodeError:
|
|||
|
|
return []
|
|||
|
|
if not isinstance(data, list):
|
|||
|
|
return []
|
|||
|
|
out: List[int] = []
|
|||
|
|
for item in data:
|
|||
|
|
try:
|
|||
|
|
out.append(int(item))
|
|||
|
|
except (TypeError, ValueError):
|
|||
|
|
continue
|
|||
|
|
return out
|