from uuid import NAMESPACE_URL, uuid5 from backend.database.schemas import Chunk from backend.embeddings.embedder import embedder from backend.core.config import settings def cosine(a: list[float], b: list[float]) -> float: return sum(x * y for x, y in zip(a, b)) class InMemoryVectorStore: def __init__(self) -> None: self._vectors: dict[str, list[tuple[Chunk, list[float]]]] = {} def index(self, repo_id: str, chunks: list[Chunk]) -> None: self._vectors[repo_id] = [(chunk, embedder.embed(chunk.content)) for chunk in chunks] def search(self, repo_id: str, query: str, top_k: int = 8) -> list[tuple[Chunk, float]]: query_vector = embedder.embed(query) scored = [ (chunk, cosine(query_vector, vector)) for chunk, vector in self._vectors.get(repo_id, []) ] return sorted(scored, key=lambda item: item[1], reverse=True)[:top_k] class QdrantVectorStore: def __init__(self) -> None: if not settings.qdrant_url: raise ValueError("Qdrant URL is required when use_qdrant=True") try: from qdrant_client import QdrantClient from qdrant_client.http import models as rest except ImportError as exc: raise ImportError( "qdrant-client is required for Qdrant vector storage. Install qdrant-client." ) from exc self.rest = rest self.client = QdrantClient( url=settings.qdrant_url, api_key=settings.qdrant_api_key, prefer_grpc=False, timeout=settings.qdrant_timeout_seconds, ) self.collection_name = settings.qdrant_collection self._ensure_collection() def _ensure_collection(self) -> None: try: collections = self.client.get_collections().collections existing = [collection.name for collection in collections] except Exception: existing = [] if self.collection_name not in existing: self.client.recreate_collection( collection_name=self.collection_name, vectors_config=self.rest.VectorParams( size=settings.embedding_dimensions, distance=self.rest.Distance.COSINE, ), ) self._ensure_payload_indexes() def _ensure_payload_indexes(self) -> None: try: self.client.create_payload_index( collection_name=self.collection_name, field_name="repo_id", field_schema=self.rest.PayloadSchemaType.KEYWORD, ) except Exception as exc: message = str(exc).lower() if "already exists" not in message and "conflict" not in message: raise def _point_id(self, chunk_id: str) -> str: return str(uuid5(NAMESPACE_URL, chunk_id)) def index(self, repo_id: str, chunks: list[Chunk]) -> None: points = [] for chunk in chunks: points.append( self.rest.PointStruct( id=self._point_id(chunk.id), vector=embedder.embed(chunk.content), payload={ "chunk_id": chunk.id, "repo_id": repo_id, "path": chunk.path, "language": chunk.language, "symbol": chunk.symbol, "kind": chunk.kind, "start_line": chunk.start_line, "end_line": chunk.end_line, "content": chunk.content, }, ) ) batch_size = max(settings.qdrant_upsert_batch_size, 1) for start in range(0, len(points), batch_size): batch = points[start : start + batch_size] self.client.upsert(collection_name=self.collection_name, points=batch) def search(self, repo_id: str, query: str, top_k: int = 8) -> list[tuple[Chunk, float]]: query_vector = embedder.embed(query) query_filter = self.rest.Filter( must=[ self.rest.FieldCondition( key="repo_id", match=self.rest.MatchValue(value=repo_id), ) ] ) if hasattr(self.client, "search"): response = self.client.search( collection_name=self.collection_name, query_vector=query_vector, limit=top_k, query_filter=query_filter, ) else: query_response = self.client.query_points( collection_name=self.collection_name, query=query_vector, limit=top_k, query_filter=query_filter, ) response = query_response.points results: list[tuple[Chunk, float]] = [] for point in response: payload = point.payload or {} chunk = Chunk( id=payload.get("chunk_id", str(point.id)), repo_id=repo_id, path=payload.get("path", ""), language=payload.get("language", ""), symbol=payload.get("symbol", "file"), kind=payload.get("kind", "text"), start_line=int(payload.get("start_line", 1)), end_line=int(payload.get("end_line", 1)), content=payload.get("content", ""), ) score = float(point.score or 0.0) results.append((chunk, score)) return results vector_store = QdrantVectorStore() if settings.use_qdrant else InMemoryVectorStore()