gitrag-free-mcp / packages /core /src /Retriever.py
jacklalane's picture
Stabilize MCP HTTP transport
9141276 verified
Raw
History Blame Contribute Delete
6.6 kB
"""Hybrid retrieval orchestration with optional Qwen3 reranking."""
from __future__ import annotations
import logging
from dataclasses import dataclass
from threading import Lock
from typing import Any, List, Optional, Protocol, Sequence
import numpy as np
from sentence_transformers import CrossEncoder
from Calculators.EmbeddingCalculator import EmbeddingCalculator
from Chunker.Chunk import Chunk
from ComponentLoader import load_components
from Persistence.Persist import PersistenceAdapter
from constants import (
DEFAULT_INITIAL_RETRIEVAL_LIMIT,
MAX_INITIAL_RETRIEVAL_LIMIT,
DEFAULT_RERANK_TASK_INSTRUCTION,
DEFAULT_RERANKER_MODEL,
RERANK_BATCH_SIZE,
RETRIEVAL_QUERY_PREFIX, EMBEDDING_BACKEND,
)
LOG4J_FORMAT = "%(asctime)s %(levelname)-5s %(name)s - %(message)s"
DATE_FORMAT = "%Y-%m-%d %H:%M:%S"
logging.basicConfig(
level=logging.INFO,
format=LOG4J_FORMAT,
datefmt=DATE_FORMAT
)
logger = logging.getLogger("indexer")
class Reranker:
"""Reranker backed by `Qwen/Qwen3-Reranker-0.6B`.
The heavy model objects are cached at class scope so multiple retrievers
can share one loaded model instance.
"""
_model: Any = None
_lock = Lock()
_inference_lock = Lock()
_loaded_model_name: str = ""
def __init__(
self,
model_name: str = DEFAULT_RERANKER_MODEL,
task_instruction: str | None = None
) -> None:
self._model_name = model_name
self._task_instruction = task_instruction or DEFAULT_RERANK_TASK_INSTRUCTION
self._batch_size = max(1, RERANK_BATCH_SIZE)
self._ensure_loaded(model_name=model_name)
@classmethod
def _ensure_loaded(cls, model_name: str = DEFAULT_RERANKER_MODEL) -> None:
if cls._model is not None and cls._loaded_model_name == model_name:
return
with cls._lock:
if cls._model is not None and cls._loaded_model_name == model_name:
return
try:
logger.info("Loading reranker model from %s with backend %s", model_name, EMBEDDING_BACKEND)
cls._model = CrossEncoder(
model_name,
trust_remote_code=True,
device="cpu",
backend=EMBEDDING_BACKEND,
)
cls._loaded_model_name = model_name
except Exception as e:
logger.warning("ONNX/OPENVINO backend failed for reranker, falling back to PyTorch: %r", e)
cls._model = CrossEncoder(model_name, trust_remote_code=True, device="cpu")
cls._loaded_model_name = model_name
def score(self, query: str, candidates: Sequence[Chunk]) -> List[float]:
if not candidates:
return []
pairs = [
(f"Instruct: {self._task_instruction}\nQuery: {query}", c.chunk)
for c in candidates
]
with self.__class__._inference_lock:
scores = self.__class__._model.predict(
pairs,
batch_size=self._batch_size,
show_progress_bar=False,
)
return [float(s) for s in scores]
@dataclass
class Retriever:
persistence: PersistenceAdapter
embedding_calculator: EmbeddingCalculator
reranker: Optional[Reranker] = None
initial_limit: int = DEFAULT_INITIAL_RETRIEVAL_LIMIT
_loaded_model_name: str = ""
def retrieve(self, query: str, *, top_k: int = 10, repo: str | None = None, branch: str | None = None, ) \
-> List[Chunk]:
normalized_query = (query or "").strip()
if not normalized_query:
return []
query_embedding = self._calculate_query_embedding(RETRIEVAL_QUERY_PREFIX + normalized_query)
initial_limit = max(top_k, min(max(self.initial_limit, top_k), MAX_INITIAL_RETRIEVAL_LIMIT))
candidates = self.persistence.search(
query_embedding=query_embedding,
query_text=normalized_query,
limit=initial_limit,
repo=repo,
branch=branch,
)
if not candidates:
return []
if self.reranker is None:
return candidates[:top_k]
scores = self.reranker.score(normalized_query, candidates)
if len(scores) != len(candidates):
raise ValueError(f"Reranker returned mismatched score length: {len(scores)} vs {len(candidates)}")
ranked = sorted(zip(scores, candidates), key=lambda item: item[0], reverse=True)
return [chunk for _, chunk in ranked[:top_k]]
def _calculate_query_embedding(self, query: str) -> np.ndarray:
raw = self.embedding_calculator.calculate(query)
if isinstance(raw, (bytearray, memoryview)):
raw = bytes(raw)
if not isinstance(raw, bytes):
raise TypeError(f"EmbeddingCalculator.calculate must return bytes, got {type(raw)}")
return np.frombuffer(raw, dtype=np.float32)
def main():
import argparse, json, sys
parser = argparse.ArgumentParser(description="GitRag Retriever")
parser.add_argument("query", nargs="?", help="Search query")
parser.add_argument("--top-k", type=int, default=10)
parser.add_argument("--repo", default=None)
parser.add_argument("--branch", default=None)
parser.add_argument("--server", action="store_true")
parser.add_argument("--port", type=int, default=8000)
args = parser.parse_args()
if not args.server and not args.query:
parser.print_help()
sys.exit(1)
calc, persist = load_components()
reranker = Reranker()
retriever = Retriever(persist, calc, reranker=reranker)
if args.server:
from gitrag_mcp_server.server import create_mcp_server
mcp = create_mcp_server(retriever=retriever)
mcp.run(transport="http", port=args.port, host="0.0.0.0", stateless_http=True, json_response=True, show_banner=False)
else:
results = retriever.retrieve(args.query, top_k=args.top_k, repo=args.repo, branch=args.branch)
output = [
{
"chunk": c.chunk,
"path": c.path,
"repo": c.repo,
"branch": c.branch,
"language": c.language,
"start_rc": c.start_rc,
"end_rc": c.end_rc,
"start_bytes": c.start_bytes,
"end_bytes": c.end_bytes,
"signature": c.signature,
}
for c in results
]
sys.stdout.write(json.dumps(output))
if __name__ == "__main__":
main()