Thicket/thicket/qdrant_store.py

198 lines
7.5 KiB
Python

"""Qdrant vector store — the service-backed target in the registry.
Invariants:
* ``replace_document`` deletes every existing point for the document
(filter on ``doc_key``) before upserting, so re-ingesting an edited
or shrunken source replaces its vectors exactly — the point count
after N re-ingests equals the count after the first.
* Upserts travel in batches of 64 — one document never arrives as a
single unbounded request.
* Point IDs are deterministic (MD5-UUID of doc_key|index), making
re-ingestion an overwrite by construction.
* A model/collection dimension mismatch is a hard error with an
actionable message; vectors are never silently mixed.
Step-down chain: ``search`` uses ``query_points`` (qdrant-client >= 1.10)
and steps down to ``search`` on older clients.
"""
from __future__ import annotations
import hashlib
import uuid
from collections.abc import Callable
from .chunker import Chunk
from .vector_stores import VectorStoreError, _payload, _snippet_payload, doc_key
UPSERT_BATCH = 64
def _point_id(key: str, idx: int) -> str:
raw = f"{key}|{idx}".encode()
return str(uuid.UUID(bytes=hashlib.md5(raw).digest()))
class QdrantStore:
"""Handles embedding indexing into a Qdrant collection. Embeddings
are computed by the attached EmbeddingEngine; the store never
embeds on its own."""
def __init__(self, host: str, port: int, collection: str, dim: int,
log: Callable[[str], None] = lambda _msg: None):
self.host = host
self.port = port
self.collection = collection
self.dim = dim
self.model_hint = "selected model"
self._engine = None
self._log = log
self._client = None
def client(self):
"""Lazy client — imports qdrant_client on first use so the GUI
never needs it installed to start."""
if self._client is None:
try:
from qdrant_client import QdrantClient
except ImportError as e:
raise VectorStoreError(
"qdrant-client not installed — run: pip install 'thicket[ingest]'"
) from e
self._client = QdrantClient(
url=f"http://{self.host}:{self.port}", timeout=10,
check_compatibility=False, # API gaps handled by step-down
)
return self._client
def ensure_collection(self) -> None:
"""Create the collection when missing; verify dimensions when it
exists."""
from qdrant_client.models import Distance, VectorParams
client = self.client()
try:
names = [c.name for c in client.get_collections().collections]
except Exception as e:
raise VectorStoreError(
f"cannot reach Qdrant at {self.host}:{self.port} — {e}"
) from e
if self.collection not in names:
self._log(f"Creating collection '{self.collection}' (dim={self.dim})...")
client.create_collection(
collection_name=self.collection,
vectors_config=VectorParams(size=self.dim, distance=Distance.COSINE),
)
self._index_payload(client)
return
self._index_payload(client)
existing_dim = self._collection_dim(client)
if isinstance(existing_dim, int) and self.dim and existing_dim != self.dim:
raise VectorStoreError(
f"collection '{self.collection}' has {existing_dim}-dim vectors "
f"but '{self.model_hint}' produces {self.dim} — pick the "
f"matching embedding model or a new collection name"
)
def _index_payload(self, client) -> None:
"""Keyword index on doc_key — every re-ingest deletes by that
filter, and unindexed payload filters scan the whole collection."""
from qdrant_client.models import PayloadSchemaType
client.create_payload_index(
collection_name=self.collection,
field_name="doc_key",
field_schema=PayloadSchemaType.KEYWORD,
)
def _collection_dim(self, client) -> int | None:
"""Vector size of the existing collection; None when the server
does not report it."""
vectors = client.get_collection(self.collection).config.params.vectors
size = getattr(vectors, "size", None)
return size if isinstance(size, int) else None
def replace_document(self, title: str, obsidian_rel_path: str,
chunks: list[Chunk]) -> int:
"""Index every chunk of a document, replacing any previous
version. Returns the number of points written."""
from qdrant_client.models import FieldCondition, Filter, MatchValue, PointStruct
if not chunks:
return 0
client = self.client()
key = doc_key(title, obsidian_rel_path)
# Exact replacement: drop the previous version before writing,
# so a shrunken source leaves no stale tail chunks behind. The
# should-clause also matches pre-1.2 points (which lack
# doc_key) via their title+path — upgrading installations clean
# themselves on first re-ingest.
client.delete(
collection_name=self.collection,
points_selector=Filter(should=[
FieldCondition(key="doc_key", match=MatchValue(value=key)),
Filter(must=[
FieldCondition(key="document_title",
match=MatchValue(value=title)),
FieldCondition(key="obsidian_path",
match=MatchValue(value=obsidian_rel_path)),
]),
]),
)
vectors = self._engine.embed([c.contextual_text for c in chunks])
points = [
PointStruct(
id=_point_id(key, idx),
vector=vector,
payload=_payload(title, obsidian_rel_path, chunk, idx),
)
for idx, (chunk, vector) in enumerate(zip(chunks, vectors))
]
batches = (points[i:i + UPSERT_BATCH]
for i in range(0, len(points), UPSERT_BATCH))
for batch in batches:
client.upsert(collection_name=self.collection, points=batch)
return len(points)
def set_embedder(self, engine) -> None:
"""Attach the EmbeddingEngine; records its model name for error
messages and adopts its dimension when none was given."""
self._engine = engine
self.model_hint = engine.model_name
if engine.dim and not self.dim:
self.dim = engine.dim
def close(self) -> None:
if self._client is not None:
self._client.close()
self._client = None
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
"""Semantic search. Returns [{score, payload}] best-first."""
client = self.client()
try:
hits = client.query_points(
collection_name=self.collection,
query=vector,
limit=limit,
with_payload=True,
).points
except AttributeError:
# Step down: qdrant-client < 1.10 exposes search() only.
hits = client.search(
collection_name=self.collection,
query_vector=vector,
limit=limit,
with_payload=True,
)
return [{"score": h.score,
"payload": _snippet_payload(h.payload or {})}
for h in hits]