ai / backend /embeddings /vector_store.py
3v324v23's picture
agent
ee9a09c
Raw
History Blame Contribute Delete
5.75 kB
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()