| 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() |
|
|