transcription/src/rag/engine/rerank.py

106 lines
3.2 KiB
Python
Raw Normal View History

"""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