1036 lines
44 KiB
Python
1036 lines
44 KiB
Python
"""Vector store abstraction — one protocol, five open-source targets.
|
|
|
|
Every target implements the same contract (``ensure_collection``,
|
|
``set_embedder``, ``replace_document``, ``search``) with the same
|
|
payload schema, so the pipeline, retrieval strip, and probe are target-
|
|
agnostic. Heavy client libraries import lazily inside each store — a
|
|
missing target is a readable error at point of use, never a crash at
|
|
startup.
|
|
|
|
Targets:
|
|
|
|
========== ============ ==============================================
|
|
key mode library
|
|
========== ============ ==============================================
|
|
qdrant service qdrant-client (HTTP to a Qdrant server)
|
|
chroma embedded chromadb (PersistentClient, on-disk)
|
|
lancedb embedded lancedb (columnar, on-disk)
|
|
faiss file faiss-cpu (IndexIDMap2 + JSON sidecar)
|
|
milvus embedded pymilvus (Milvus Lite local database)
|
|
========== ============ ==============================================
|
|
|
|
Embedded/file targets keep their data under ``<data_dir>`` (the caller
|
|
passes ``<vault>/.thicket/<target>``), so a vault stays a single
|
|
portable tree.
|
|
|
|
Payload contract (identical across targets):
|
|
|
|
document_title, obsidian_path, section_header, content,
|
|
chunk_index, doc_key
|
|
|
|
``doc_key`` (MD5 of ``title|rel_path``) is the deletion key for
|
|
exact-replacement re-ingest — a collision-free handle that keeps
|
|
filter expressions free of user-controlled strings.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import json
|
|
import re
|
|
from collections.abc import Callable
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
|
|
from .chunker import Chunk
|
|
|
|
UPSERT_BATCH = 64
|
|
|
|
|
|
class VectorStoreError(Exception):
|
|
"""Raised for target availability, connectivity, or write failures."""
|
|
|
|
|
|
def doc_key(title: str, rel_path: str) -> str:
|
|
"""Stable per-document key used for exact-replacement deletes."""
|
|
return hashlib.md5(f"{title}|{rel_path}".encode()).hexdigest()
|
|
|
|
|
|
def _payload(title: str, rel_path: str, chunk: Chunk, idx: int) -> dict:
|
|
payload = {
|
|
"document_title": title,
|
|
"obsidian_path": rel_path,
|
|
"section_header": chunk.header,
|
|
"content": chunk.text,
|
|
"chunk_index": idx,
|
|
"chunk_kind": getattr(chunk, "kind", "prose"),
|
|
"doc_key": doc_key(title, rel_path),
|
|
}
|
|
if getattr(chunk, "lang", None):
|
|
payload["lang"] = chunk.lang
|
|
if getattr(chunk, "source", None):
|
|
# Origin file within the input tree — programming RAG answers
|
|
# "which file is this from", not just "which document title".
|
|
payload["source_path"] = chunk.source
|
|
return payload
|
|
|
|
|
|
def _snippet_payload(payload: dict) -> dict:
|
|
"""Payload fields the retrieval strip displays (drop internal keys)."""
|
|
return {k: v for k, v in payload.items() if k != "doc_key"}
|
|
|
|
|
|
class BaseVectorStore:
|
|
"""Shared behavior for embedded/file targets: embedder attachment,
|
|
dimension bookkeeping, batched writes. Subclasses own the client
|
|
lifecycle and implement the four protocol methods."""
|
|
|
|
def __init__(self, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
self.collection = collection
|
|
self.dim = dim
|
|
self.model_hint = "selected model"
|
|
self._engine = None
|
|
self._log = log
|
|
|
|
def set_embedder(self, engine) -> None:
|
|
"""Attach the EmbeddingEngine; adopts its dimension when none
|
|
was declared."""
|
|
self._engine = engine
|
|
self.model_hint = engine.model_name
|
|
if engine.dim and not self.dim:
|
|
self.dim = engine.dim
|
|
|
|
def _require_dimension(self) -> None:
|
|
if not self.dim:
|
|
raise VectorStoreError(
|
|
f"'{self.model_hint}' has no catalog dimension — the "
|
|
f"{type(self).__name__} target requires one"
|
|
)
|
|
|
|
def _vectors(self, chunks: list[Chunk]) -> list[list[float]]:
|
|
return self._engine.embed([c.contextual_text for c in chunks])
|
|
|
|
def close(self) -> None:
|
|
"""Release connections and file locks. Embedded file targets
|
|
(Milvus Lite in particular) hold exclusive locks while
|
|
connected — release or the next process cannot open them."""
|
|
for attr in ("_client", "_conn", "_db"):
|
|
obj = getattr(self, attr, None)
|
|
if obj is None:
|
|
continue
|
|
closer = getattr(obj, "close", None)
|
|
if closer is not None:
|
|
try:
|
|
closer()
|
|
except Exception:
|
|
pass
|
|
setattr(self, attr, None)
|
|
# Milvus Lite: closing the client leaves the embedded server
|
|
# thread holding the database file lock — release it too.
|
|
try:
|
|
from milvus_lite.server_manager import server_manager_instance
|
|
server_manager_instance.release_all()
|
|
except Exception:
|
|
pass
|
|
|
|
def _require_module(self, module: str) -> None:
|
|
try:
|
|
return __import__(module)
|
|
except ImportError as e:
|
|
raise VectorStoreError(
|
|
f"{module} not installed — run: pip install 'thicket[{module}]' "
|
|
f"or pick another vector target"
|
|
) from e
|
|
|
|
# protocol: ensure_collection / replace_document / search
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# Chroma (embedded)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
_CHROMA_NAME_OK = (
|
|
lambda name: 3 <= len(name) <= 512
|
|
and all(c.isalnum() or c in "._-" for c in name)
|
|
and name[0].isalnum() and name[-1].isalnum()
|
|
)
|
|
|
|
|
|
class ChromaStore(BaseVectorStore):
|
|
"""Chroma persistent collection; cosine space, on-disk under the
|
|
data dir. Dimensions are implicit — Chroma validates them at write."""
|
|
|
|
def __init__(self, data_dir: Path, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
super().__init__(collection, dim, log)
|
|
self._require_module("chromadb")
|
|
import chromadb
|
|
self._client = chromadb.PersistentClient(path=str(data_dir))
|
|
self._col = None
|
|
|
|
def ensure_collection(self) -> None:
|
|
if not _CHROMA_NAME_OK(self.collection):
|
|
raise VectorStoreError(
|
|
f"Chroma collection names need 3-512 chars of [a-zA-Z0-9._-] "
|
|
f"(got '{self.collection}')"
|
|
)
|
|
self._col = self._client.get_or_create_collection(
|
|
name=self.collection, metadata={"hnsw:space": "cosine"}
|
|
)
|
|
|
|
def _collection(self):
|
|
if self._col is None:
|
|
self.ensure_collection()
|
|
return self._col
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
import uuid as _uuid
|
|
|
|
if not chunks:
|
|
return 0
|
|
col = self._collection()
|
|
key = doc_key(title, obsidian_rel_path)
|
|
col.delete(where={"doc_key": {"$eq": key}})
|
|
|
|
vectors = self._vectors(chunks)
|
|
for i in range(0, len(chunks), UPSERT_BATCH):
|
|
batch = chunks[i:i + UPSERT_BATCH]
|
|
col.add(
|
|
ids=[str(_uuid.UUID(bytes=hashlib.md5(
|
|
f"{key}|{i + j}".encode()).digest())) for j in range(len(batch))],
|
|
embeddings=vectors[i:i + UPSERT_BATCH],
|
|
documents=[c.text for c in batch],
|
|
metadatas=[_payload(title, obsidian_rel_path, c, i + j)
|
|
for j, c in enumerate(batch)],
|
|
)
|
|
return len(chunks)
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
result = self._collection().query(
|
|
query_embeddings=[vector], n_results=limit,
|
|
include=["metadatas", "distances"],
|
|
)
|
|
metas = result["metadatas"][0]
|
|
dists = result["distances"][0]
|
|
return [{"score": 1.0 - d, "payload": _snippet_payload(m)}
|
|
for m, d in zip(metas, dists)]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# LanceDB (embedded)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
class LanceStore(BaseVectorStore):
|
|
"""LanceDB table; vectors are L2-normalized on write and query so
|
|
the default L2 metric ranks identically to cosine."""
|
|
|
|
def __init__(self, data_dir: Path, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
super().__init__(collection, dim, log)
|
|
self._require_module("lancedb")
|
|
import lancedb
|
|
self._db = lancedb.connect(str(data_dir))
|
|
self._data_dir = data_dir
|
|
|
|
def _table(self):
|
|
try:
|
|
return self._db.open_table(self.collection)
|
|
except Exception:
|
|
return None
|
|
|
|
def ensure_collection(self) -> None:
|
|
# Tables are created with their first data batch (the schema is
|
|
# data-derived); a missing table simply means "nothing ingested".
|
|
pass
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
if not chunks:
|
|
return 0
|
|
self._require_dimension()
|
|
import numpy as np
|
|
|
|
key = doc_key(title, obsidian_rel_path)
|
|
table = self._table()
|
|
if table is not None:
|
|
table.delete(f"doc_key = '{key}'")
|
|
|
|
# Payload is a JSON string column: the chunk schema carries
|
|
# optional fields (lang), and struct inference would break on
|
|
# their presence/absence across batches.
|
|
records = [
|
|
{"doc_key": key,
|
|
"vector": (np.asarray(v) / np.linalg.norm(v)).tolist(),
|
|
"payload": json.dumps(_payload(title, obsidian_rel_path, c, i))}
|
|
for i, (c, v) in enumerate(zip(chunks, self._vectors(chunks)))
|
|
]
|
|
if table is None:
|
|
self._db.create_table(self.collection, data=records)
|
|
else:
|
|
table.add(records)
|
|
return len(chunks)
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
import numpy as np
|
|
|
|
table = self._table()
|
|
if table is None:
|
|
raise VectorStoreError(
|
|
f"collection '{self.collection}' is empty — ingest first"
|
|
)
|
|
q = np.asarray(vector, dtype="float32")
|
|
q = (q / np.linalg.norm(q)).tolist()
|
|
rows = table.search(q).limit(limit).to_list()
|
|
# Lance distances are squared L2; normalized squared L2 d
|
|
# satisfies cos = 1 - d/2 — matching every other target's
|
|
# cosine score.
|
|
return [{"score": 1.0 - r.get("_distance", 0.0) / 2.0,
|
|
"payload": _snippet_payload(json.loads(r["payload"]))}
|
|
for r in rows]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# FAISS (file)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
class FaissStore(BaseVectorStore):
|
|
"""FAISS IndexIDMap2 over an inner-product index with L2-normalized
|
|
vectors (= cosine). Payloads live in a JSON sidecar keyed by the
|
|
same deterministic ids as the index."""
|
|
|
|
def __init__(self, data_dir: Path, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
super().__init__(collection, dim, log)
|
|
faiss = self._require_module("faiss")
|
|
self._faiss = faiss
|
|
self._index_path = data_dir / f"{collection}.faiss"
|
|
self._meta_path = data_dir / f"{collection}.json"
|
|
self._index = None
|
|
self._meta: dict[str, dict] = {}
|
|
|
|
def ensure_collection(self) -> None:
|
|
self._require_dimension()
|
|
if self._index is not None:
|
|
return
|
|
if self._index_path.exists():
|
|
self._index = self._faiss.read_index(str(self._index_path))
|
|
self._meta = json.loads(self._meta_path.read_text(encoding="utf-8"))
|
|
else:
|
|
self._index = self._faiss.IndexIDMap2(
|
|
self._faiss.IndexFlatIP(self.dim)
|
|
)
|
|
self._persist()
|
|
|
|
def _persist(self) -> None:
|
|
self._index_path.parent.mkdir(parents=True, exist_ok=True)
|
|
self._faiss.write_index(self._index, str(self._index_path))
|
|
self._meta_path.write_text(
|
|
json.dumps(self._meta, ensure_ascii=False), encoding="utf-8"
|
|
)
|
|
|
|
@staticmethod
|
|
def _faiss_id(key: str, idx: int) -> int:
|
|
digest = hashlib.md5(f"{key}|{idx}".encode()).digest()
|
|
return int.from_bytes(digest[:8], "big") & (2**63 - 1)
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
import numpy as np
|
|
|
|
if not chunks:
|
|
return 0
|
|
self.ensure_collection()
|
|
key = doc_key(title, obsidian_rel_path)
|
|
|
|
stale = [pid for pid, p in self._meta.items()
|
|
if p.get("doc_key") == key]
|
|
if stale:
|
|
# Old chunk ids are (key, 0..count-1) by construction.
|
|
self._index.remove_ids(np.array(
|
|
[self._faiss_id(key, i) for i in range(len(stale))],
|
|
dtype="int64"))
|
|
for s in stale:
|
|
del self._meta[s]
|
|
|
|
vectors = np.asarray(self._vectors(chunks), dtype="float32")
|
|
vectors /= np.linalg.norm(vectors, axis=1, keepdims=True)
|
|
ids = np.array([self._faiss_id(key, i) for i in range(len(chunks))],
|
|
dtype="int64")
|
|
self._index.add_with_ids(vectors, ids)
|
|
for i, chunk in enumerate(chunks):
|
|
pid = str(ids[i])
|
|
self._meta[pid] = _payload(title, obsidian_rel_path, chunk, i)
|
|
self._persist()
|
|
return len(chunks)
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
import numpy as np
|
|
|
|
self.ensure_collection()
|
|
q = np.asarray([vector], dtype="float32")
|
|
q /= np.linalg.norm(q)
|
|
scores, ids = self._index.search(q, limit)
|
|
return [{"score": float(s), "payload": _snippet_payload(self._meta[str(i)])}
|
|
for s, i in zip(scores[0], ids[0]) if i != -1]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# Milvus Lite (embedded)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
class MilvusStore(BaseVectorStore):
|
|
"""Milvus Lite local database via pymilvus MilvusClient — one
|
|
flat cosine index, JSON payloads, VARCHAR primary ids."""
|
|
|
|
def __init__(self, data_dir: Path, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
super().__init__(collection, dim, log)
|
|
self._require_module("pymilvus")
|
|
from pymilvus import MilvusClient
|
|
data_dir.mkdir(parents=True, exist_ok=True)
|
|
import time
|
|
|
|
last_error: Exception | None = None
|
|
for attempt in range(3): # lock may linger one beat after close
|
|
try:
|
|
self._client = MilvusClient(str(data_dir / f"{collection}.db"))
|
|
break
|
|
except Exception as e: # noqa: BLE001 — retry, then report
|
|
last_error = e
|
|
time.sleep(2)
|
|
else:
|
|
raise VectorStoreError(
|
|
f"Milvus Lite unavailable: {last_error} — another process "
|
|
f"may hold the database lock, or run: "
|
|
f"pip install 'pymilvus[milvus_lite]'"
|
|
) from last_error
|
|
|
|
def ensure_collection(self) -> None:
|
|
self._require_dimension()
|
|
if self._client.has_collection(self.collection):
|
|
self._client.load_collection(self.collection) # no-op when loaded
|
|
return
|
|
from pymilvus import DataType
|
|
|
|
schema = self._client.create_schema(auto_id=False)
|
|
schema.add_field("id", DataType.VARCHAR, max_length=64, is_primary=True)
|
|
schema.add_field("vector", DataType.FLOAT_VECTOR, dim=self.dim)
|
|
schema.add_field("doc_key", DataType.VARCHAR, max_length=64)
|
|
schema.add_field("payload", DataType.JSON)
|
|
index = self._client.prepare_index_params()
|
|
index.add_index(field_name="vector", index_type="FLAT",
|
|
metric_type="COSINE")
|
|
self._client.create_collection(self.collection, schema=schema,
|
|
index_params=index)
|
|
self._client.load_collection(self.collection)
|
|
self._log(f"Creating collection '{self.collection}' in Milvus Lite...")
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
import uuid as _uuid
|
|
|
|
if not chunks:
|
|
return 0
|
|
self.ensure_collection()
|
|
key = doc_key(title, obsidian_rel_path)
|
|
self._client.delete(self.collection,
|
|
filter=f'doc_key == "{key}"')
|
|
|
|
rows = [
|
|
{"id": str(_uuid.UUID(bytes=hashlib.md5(
|
|
f"{key}|{i}".encode()).digest())),
|
|
"vector": v,
|
|
"doc_key": key,
|
|
"payload": _payload(title, obsidian_rel_path, c, i)}
|
|
for i, (c, v) in enumerate(zip(chunks, self._vectors(chunks)))
|
|
]
|
|
for i in range(0, len(rows), UPSERT_BATCH):
|
|
self._client.insert(self.collection, rows[i:i + UPSERT_BATCH])
|
|
# Lite has no Strong-consistency knob in this client version;
|
|
# flush makes delete+re-insert immediately visible to search.
|
|
self._client.flush(self.collection)
|
|
return len(rows)
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
self.ensure_collection()
|
|
result = self._client.search(
|
|
self.collection, data=[vector], limit=limit,
|
|
output_fields=["payload"],
|
|
search_params={"metric_type": "COSINE"},
|
|
)
|
|
return [{"score": float(hit["distance"]),
|
|
"payload": _snippet_payload(hit["entity"].get("payload", {}))}
|
|
for hit in result[0]]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# Weaviate (service)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
def _weaviate_name(name: str) -> str:
|
|
"""Weaviate requires collection names to start uppercase with
|
|
[A-Za-z0-9_] bodies — map deterministically and keep the mapping
|
|
logged."""
|
|
cleaned = re.sub(r"[^A-Za-z0-9_]", "_", name)
|
|
return (cleaned[:1].upper() + cleaned[1:]) or "Thicket"
|
|
|
|
|
|
class WeaviateStore(BaseVectorStore):
|
|
"""Weaviate server (container or compose), HNSW cosine index.
|
|
Connects over HTTP (host/port) with the standard gRPC port 50051."""
|
|
|
|
GRPC_PORT = 50051
|
|
|
|
def __init__(self, host: str, port: int, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None, **_ignored):
|
|
super().__init__(collection, dim, log)
|
|
self._require_module("weaviate")
|
|
import weaviate
|
|
if host == "localhost":
|
|
# The gRPC channel resolves localhost to ::1 first, which
|
|
# host-networked servers routinely refuse — pin IPv4.
|
|
host = "127.0.0.1"
|
|
try:
|
|
self._client = weaviate.connect_to_local(
|
|
host=host, port=port, grpc_port=self.GRPC_PORT,
|
|
)
|
|
except Exception as e:
|
|
raise VectorStoreError(
|
|
f"cannot connect to Weaviate at {host}:{port} "
|
|
f"(gRPC {self.GRPC_PORT}): {e}"
|
|
) from e
|
|
self._mapped = _weaviate_name(collection)
|
|
if self._mapped != collection:
|
|
self._log(f"Weaviate maps collection '{collection}' "
|
|
f"-> '{self._mapped}' (name rules).")
|
|
|
|
def ensure_collection(self) -> None:
|
|
from weaviate.classes.config import (
|
|
Configure, DataType, Property, Tokenization, VectorDistances,
|
|
)
|
|
if not self._client.collections.exists(self._mapped):
|
|
self._log(f"Creating collection '{self._mapped}' in Weaviate...")
|
|
self._client.collections.create(
|
|
name=self._mapped,
|
|
vector_index_config=Configure.VectorIndex.hnsw(
|
|
distance_metric=VectorDistances.COSINE),
|
|
properties=[
|
|
Property(name="doc_key", data_type=DataType.TEXT,
|
|
tokenization=Tokenization.FIELD),
|
|
Property(name="document_title", data_type=DataType.TEXT),
|
|
Property(name="obsidian_path", data_type=DataType.TEXT),
|
|
Property(name="section_header", data_type=DataType.TEXT),
|
|
Property(name="content", data_type=DataType.TEXT),
|
|
Property(name="chunk_index", data_type=DataType.INT),
|
|
],
|
|
)
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
if not chunks:
|
|
return 0
|
|
import uuid as _uuid
|
|
from weaviate.classes.query import Filter
|
|
|
|
coll = self._client.collections.get(self._mapped)
|
|
key = doc_key(title, obsidian_rel_path)
|
|
coll.data.delete_many(where=Filter.by_property("doc_key").equal(key))
|
|
|
|
with coll.batch.fixed_size(batch_size=UPSERT_BATCH) as batch:
|
|
for i, (chunk, vector) in enumerate(
|
|
zip(chunks, self._vectors(chunks))):
|
|
batch.add_object(
|
|
uuid=_uuid.UUID(bytes=hashlib.md5(
|
|
f"{key}|{i}".encode()).digest()),
|
|
vector=vector,
|
|
properties=_payload(title, obsidian_rel_path, chunk, i),
|
|
)
|
|
failed = coll.batch.failed_objects
|
|
if failed:
|
|
raise VectorStoreError(
|
|
f"Weaviate rejected {len(failed)} object(s): "
|
|
f"{failed[0].message if failed else '?'}"
|
|
)
|
|
return len(chunks)
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
from weaviate.classes.query import MetadataQuery
|
|
|
|
coll = self._client.collections.get(self._mapped)
|
|
response = coll.query.near_vector(
|
|
near_vector=vector, limit=limit,
|
|
return_metadata=MetadataQuery(distance=True),
|
|
)
|
|
return [{"score": 1.0 - (obj.metadata.distance or 0.0),
|
|
"payload": _snippet_payload(dict(obj.properties))}
|
|
for obj in response.objects]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# pgvector (Postgres service)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
class PgVectorStore(BaseVectorStore):
|
|
"""Postgres + pgvector: one table per collection, HNSW cosine
|
|
index, JSONB payloads. Connection comes from the standard libpq
|
|
environment (PGHOST / PGPORT / PGDATABASE / PGUSER / PGPASSWORD,
|
|
or a full PGDSN) — Unix convention, no extra config surface."""
|
|
|
|
def __init__(self, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None, **_ignored):
|
|
super().__init__(collection, dim, log)
|
|
self._require_module("psycopg")
|
|
import os
|
|
|
|
import psycopg
|
|
dsn = os.environ.get("PGDSN")
|
|
try:
|
|
self._conn = psycopg.connect(dsn or "", connect_timeout=5)
|
|
except Exception as e:
|
|
raise VectorStoreError(
|
|
f"cannot connect to Postgres ({dsn or 'libpq env'}): {e} — "
|
|
f"set PGHOST/PGUSER/PGPASSWORD/PGDATABASE or PGDSN"
|
|
) from e
|
|
try:
|
|
from pgvector.psycopg import register_vector
|
|
register_vector(self._conn)
|
|
except ImportError as e:
|
|
raise VectorStoreError(
|
|
"pgvector package not installed — run: "
|
|
"pip install 'thicket[pgvector]'"
|
|
) from e
|
|
|
|
def _table_sql(self, template: str):
|
|
"""Compose SQL with the collection name as a quoted identifier —
|
|
user-supplied names never interpolate as raw SQL."""
|
|
from psycopg import sql
|
|
return sql.SQL(template).format(tbl=sql.Identifier(self.collection))
|
|
|
|
def ensure_collection(self) -> None:
|
|
self._require_dimension()
|
|
from psycopg import sql
|
|
|
|
try:
|
|
with self._conn.cursor() as cur:
|
|
cur.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
|
cur.execute(sql.SQL(
|
|
"CREATE TABLE IF NOT EXISTS {tbl} ("
|
|
"id TEXT PRIMARY KEY, doc_key TEXT NOT NULL, "
|
|
"embedding vector({dim}), payload JSONB)"
|
|
).format(tbl=sql.Identifier(self.collection),
|
|
dim=sql.Literal(self.dim)))
|
|
cur.execute(sql.SQL(
|
|
"CREATE INDEX IF NOT EXISTS {idx} ON {tbl} USING hnsw "
|
|
"(embedding vector_cosine_ops)"
|
|
).format(idx=sql.Identifier(f"ix_{self.collection}_cos"),
|
|
tbl=sql.Identifier(self.collection)))
|
|
cur.execute(sql.SQL(
|
|
"CREATE INDEX IF NOT EXISTS {idx} ON {tbl} (doc_key)"
|
|
).format(idx=sql.Identifier(f"ix_{self.collection}_dk"),
|
|
tbl=sql.Identifier(self.collection)))
|
|
self._conn.commit()
|
|
except Exception as e:
|
|
raise VectorStoreError(
|
|
f"pgvector setup failed: {e} — the server needs the "
|
|
f"pgvector extension available (image: pgvector/pgvector)"
|
|
) from e
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
import uuid as _uuid
|
|
|
|
if not chunks:
|
|
return 0
|
|
key = doc_key(title, obsidian_rel_path)
|
|
insert = self._table_sql(
|
|
"INSERT INTO {tbl} VALUES (%s, %s, %s::vector, %s::jsonb)")
|
|
try:
|
|
with self._conn.cursor() as cur:
|
|
cur.execute(self._table_sql("DELETE FROM {tbl} WHERE doc_key = %s"),
|
|
(key,))
|
|
cur.executemany(insert, [
|
|
(str(_uuid.UUID(bytes=hashlib.md5(
|
|
f"{key}|{i}".encode()).digest())),
|
|
key, json.dumps(vector),
|
|
json.dumps(_payload(title, obsidian_rel_path, chunk, i)))
|
|
for i, (chunk, vector)
|
|
in enumerate(zip(chunks, self._vectors(chunks)))
|
|
])
|
|
self._conn.commit()
|
|
except Exception:
|
|
# Roll back so one bad document never poisons the shared
|
|
# connection for the rest of the queue.
|
|
self._conn.rollback()
|
|
raise
|
|
return len(chunks)
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
query = self._table_sql(
|
|
"SELECT payload, 1 - (embedding <=> %s::vector) AS score "
|
|
"FROM {tbl} ORDER BY embedding <=> %s::vector LIMIT %s")
|
|
qtext = json.dumps(vector)
|
|
with self._conn.cursor() as cur:
|
|
cur.execute(query, (qtext, qtext, limit))
|
|
rows = cur.fetchall()
|
|
return [{"score": float(score),
|
|
"payload": _snippet_payload(
|
|
row if isinstance(row, dict) else json.loads(row))}
|
|
for row, score in rows]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# DuckDB (embedded, VSS step-down)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
class DuckStore(BaseVectorStore):
|
|
"""DuckDB with FLOAT[dim] columns. Vectors are L2-normalized on
|
|
write and query; ranking therefore equals cosine. The vss HNSW
|
|
extension loads when available and steps down to an exact scan
|
|
otherwise (exact is fine at personal-knowledge scale)."""
|
|
|
|
def __init__(self, data_dir: Path, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
super().__init__(collection, dim, log)
|
|
self._require_module("duckdb")
|
|
import duckdb
|
|
data_dir.mkdir(parents=True, exist_ok=True)
|
|
self._db = duckdb.connect(str(data_dir / f"{collection}.duckdb"))
|
|
try:
|
|
self._db.execute("INSTALL vss")
|
|
self._db.execute("LOAD vss")
|
|
except Exception:
|
|
self._log("DuckDB vss extension unavailable — exact scan.")
|
|
|
|
def ensure_collection(self) -> None:
|
|
self._require_dimension()
|
|
self._db.execute(
|
|
f'CREATE TABLE IF NOT EXISTS "{self.collection}" ('
|
|
f'doc_key VARCHAR, embedding FLOAT[{self.dim}], payload JSON)'
|
|
)
|
|
|
|
def _normalize(self, vectors):
|
|
import numpy as np
|
|
arr = np.asarray(vectors, dtype="float32")
|
|
return (arr / np.linalg.norm(arr, axis=-1, keepdims=True).clip(1e-12)).tolist()
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
if not chunks:
|
|
return 0
|
|
self.ensure_collection()
|
|
key = doc_key(title, obsidian_rel_path)
|
|
self._db.execute(
|
|
f'DELETE FROM "{self.collection}" WHERE doc_key = ?', [key])
|
|
rows = [
|
|
(key, vector,
|
|
json.dumps(_payload(title, obsidian_rel_path, chunk, i)))
|
|
for i, (chunk, vector)
|
|
in enumerate(zip(chunks, self._normalize(self._vectors(chunks))))
|
|
]
|
|
self._db.executemany(
|
|
f'INSERT INTO "{self.collection}" '
|
|
f'VALUES (?, ?::FLOAT[{self.dim}], ?::JSON)', rows)
|
|
return len(chunks)
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
import numpy as np
|
|
|
|
self.ensure_collection()
|
|
q = np.asarray(vector, dtype="float32")
|
|
q = (q / max(float(np.linalg.norm(q)), 1e-12)).tolist()
|
|
rows = self._db.execute(
|
|
f'SELECT payload, array_distance(embedding, '
|
|
f'?::FLOAT[{self.dim}]) AS d FROM "{self.collection}" '
|
|
f'ORDER BY d LIMIT ?', [q, limit]).fetchall()
|
|
# Normalized L2 distance d satisfies cos = 1 - d²/2.
|
|
return [{"score": 1.0 - float(d) ** 2 / 2.0,
|
|
"payload": _snippet_payload(json.loads(p))}
|
|
for p, d in rows]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# sqlite-vec (embedded)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
class SqliteVecStore(BaseVectorStore):
|
|
"""sqlite-vec vec0 virtual table (float vectors, L2) with a
|
|
sidecar metadata table. Vectors L2-normalized — ranking equals
|
|
cosine; score conversion matches the other normalized-L2 targets."""
|
|
|
|
def __init__(self, data_dir: Path, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
super().__init__(collection, dim, log)
|
|
sqlite_vec = self._require_module("sqlite_vec")
|
|
import sqlite3
|
|
data_dir.mkdir(parents=True, exist_ok=True)
|
|
self._vec = sqlite_vec
|
|
self._db = sqlite3.connect(str(data_dir / f"{collection}.sqlite3"))
|
|
self._db.enable_load_extension(True)
|
|
sqlite_vec.load(self._db)
|
|
self._db.enable_load_extension(False)
|
|
|
|
def ensure_collection(self) -> None:
|
|
self._require_dimension()
|
|
self._db.execute(
|
|
f'CREATE VIRTUAL TABLE IF NOT EXISTS "vec_{self.collection}" '
|
|
f'USING vec0(embedding float[{self.dim}])')
|
|
self._db.execute(
|
|
f'CREATE TABLE IF NOT EXISTS "{self.collection}_meta" '
|
|
f'(rowid INTEGER PRIMARY KEY, doc_key TEXT, payload TEXT)')
|
|
self._db.commit()
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
if not chunks:
|
|
return 0
|
|
self.ensure_collection()
|
|
key = doc_key(title, obsidian_rel_path)
|
|
|
|
stale = self._db.execute(
|
|
f'SELECT rowid FROM "{self.collection}_meta" WHERE doc_key = ?',
|
|
(key,)).fetchall()
|
|
for (rowid,) in stale:
|
|
self._db.execute(
|
|
f'DELETE FROM "vec_{self.collection}" WHERE rowid = ?',
|
|
(rowid,))
|
|
self._db.execute(
|
|
f'DELETE FROM "{self.collection}_meta" WHERE rowid = ?',
|
|
(rowid,))
|
|
|
|
for i, (chunk, vector) in enumerate(
|
|
zip(chunks, self._normalized(self._vectors(chunks)))):
|
|
cur = self._db.execute(
|
|
f'INSERT INTO "vec_{self.collection}" (embedding) '
|
|
f'VALUES (?)', (self._vec.serialize_float32(vector),))
|
|
self._db.execute(
|
|
f'INSERT INTO "{self.collection}_meta" VALUES (?, ?, ?)',
|
|
(cur.lastrowid, key,
|
|
json.dumps(_payload(title, obsidian_rel_path, chunk, i))))
|
|
self._db.commit()
|
|
return len(chunks)
|
|
|
|
def _normalized(self, vectors: list[list[float]]) -> list[list[float]]:
|
|
import numpy as np
|
|
arr = np.asarray(vectors, dtype="float32")
|
|
norms = np.linalg.norm(arr, axis=1, keepdims=True).clip(1e-12)
|
|
return (arr / norms).tolist()
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
import numpy as np
|
|
|
|
self.ensure_collection()
|
|
q = np.asarray(vector, dtype="float32")
|
|
q = (q / max(float(np.linalg.norm(q)), 1e-12)).tolist()
|
|
rows = self._db.execute(
|
|
f'SELECT m.payload, v.distance '
|
|
f'FROM "vec_{self.collection}" v '
|
|
f'JOIN "{self.collection}_meta" m ON m.rowid = v.rowid '
|
|
f'WHERE v.embedding MATCH ? AND k = ? '
|
|
f'ORDER BY v.distance',
|
|
(self._vec.serialize_float32(q), limit)).fetchall()
|
|
return [{"score": 1.0 - float(d) ** 2 / 2.0,
|
|
"payload": _snippet_payload(json.loads(p))}
|
|
for p, d in rows]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# MariaDB (service, VECTOR columns)
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
def _mariadb_ident(name: str) -> str:
|
|
"""Table/database identifier: [A-Za-z0-9_] only — user-supplied
|
|
names never reach SQL as raw text."""
|
|
cleaned = re.sub(r"[^A-Za-z0-9_]", "_", name)
|
|
return cleaned or "thicket"
|
|
|
|
|
|
class MariaDbStore(BaseVectorStore):
|
|
"""MariaDB 11.7+ VECTOR(dim) columns with VEC_Distance_Cosine
|
|
ranking and a vector index. Connection comes from MARIADB_* env
|
|
vars (HOST/PORT/USER/PASSWORD/DATABASE, or UNIX_SOCKET).
|
|
|
|
The vector index requires primary keys <= 256 bytes; ids are
|
|
36-char UUID strings (144 bytes utf8mb4), and index creation
|
|
steps down to a logged exact scan on servers without vector
|
|
support."""
|
|
|
|
def __init__(self, collection: str, dim: int,
|
|
log: Callable[[str], None] = lambda _msg: None, **_ignored):
|
|
super().__init__(collection, dim, log)
|
|
self._require_module("pymysql")
|
|
import os
|
|
|
|
import pymysql
|
|
self._pymysql = pymysql
|
|
self._db = _mariadb_ident(os.environ.get("MARIADB_DATABASE", "thicket"))
|
|
self._table = _mariadb_ident(collection)
|
|
if self._table != collection:
|
|
self._log(f"MariaDB maps collection '{collection}' "
|
|
f"-> table '{self._table}' (identifier rules).")
|
|
try:
|
|
self._conn = pymysql.connect(
|
|
host=os.environ.get("MARIADB_HOST", "127.0.0.1"),
|
|
port=int(os.environ.get("MARIADB_PORT", "3306")),
|
|
user=os.environ.get("MARIADB_USER", "root"),
|
|
password=os.environ.get("MARIADB_PASSWORD", ""),
|
|
unix_socket=os.environ.get("MARIADB_UNIX_SOCKET") or None,
|
|
connect_timeout=5,
|
|
)
|
|
except Exception as e:
|
|
raise VectorStoreError(
|
|
f"cannot connect to MariaDB: {e} — set MARIADB_HOST / "
|
|
f"MARIADB_USER / MARIADB_PASSWORD (or MARIADB_UNIX_SOCKET)"
|
|
) from e
|
|
with self._conn.cursor() as cur:
|
|
cur.execute(f"CREATE DATABASE IF NOT EXISTS `{self._db}`")
|
|
cur.execute(f"USE `{self._db}`")
|
|
self._conn.commit()
|
|
|
|
def ensure_collection(self) -> None:
|
|
self._require_dimension()
|
|
with self._conn.cursor() as cur:
|
|
cur.execute(
|
|
f"CREATE TABLE IF NOT EXISTS `{self._db}`.`{self._table}` ("
|
|
f"id VARCHAR(36) PRIMARY KEY, "
|
|
f"doc_key VARCHAR(64) NOT NULL, "
|
|
f"embedding VECTOR({self.dim}) NOT NULL, "
|
|
f"payload JSON)")
|
|
cur.execute(
|
|
"SELECT COLUMN_TYPE FROM information_schema.COLUMNS "
|
|
"WHERE TABLE_SCHEMA = %s AND TABLE_NAME = %s "
|
|
"AND COLUMN_NAME = 'embedding'",
|
|
(self._db, self._table))
|
|
row = cur.fetchone()
|
|
if row and f"vector({self.dim})" not in str(row[0]).lower():
|
|
raise VectorStoreError(
|
|
f"table `{self._db}`.`{self._table}` was built at a "
|
|
f"different vector dimension ({row[0]}) than "
|
|
f"'{self.model_hint}' produces ({self.dim}) — pick a "
|
|
f"new collection name or drop the table")
|
|
try:
|
|
cur.execute(
|
|
f"ALTER TABLE `{self._db}`.`{self._table}` "
|
|
f"ADD VECTOR INDEX IF NOT EXISTS vix_{self._table} "
|
|
f"(embedding)")
|
|
except Exception:
|
|
self._log("MariaDB vector index unavailable — exact scan.")
|
|
cur.execute(
|
|
f"ALTER TABLE `{self._db}`.`{self._table}` "
|
|
f"ADD INDEX IF NOT EXISTS ix_{self._table}_dk (doc_key)")
|
|
self._conn.commit()
|
|
|
|
def replace_document(self, title: str, obsidian_rel_path: str,
|
|
chunks: list[Chunk]) -> int:
|
|
import uuid as _uuid
|
|
|
|
if not chunks:
|
|
return 0
|
|
self.ensure_collection()
|
|
key = doc_key(title, obsidian_rel_path)
|
|
with self._conn.cursor() as cur:
|
|
cur.execute(
|
|
f"DELETE FROM `{self._db}`.`{self._table}` "
|
|
f"WHERE doc_key = %s", (key,))
|
|
cur.executemany(
|
|
f"INSERT INTO `{self._db}`.`{self._table}` "
|
|
f"VALUES (%s, %s, VEC_FromText(%s), %s)",
|
|
[(str(_uuid.UUID(bytes=hashlib.md5(
|
|
f"{key}|{i}".encode()).digest())),
|
|
key, json.dumps(vector),
|
|
json.dumps(_payload(title, obsidian_rel_path, chunk, i)))
|
|
for i, (chunk, vector)
|
|
in enumerate(zip(chunks, self._vectors(chunks)))])
|
|
self._conn.commit()
|
|
return len(chunks)
|
|
|
|
def search(self, vector: list[float], limit: int = 5) -> list[dict]:
|
|
self.ensure_collection()
|
|
qtext = json.dumps(vector)
|
|
with self._conn.cursor() as cur:
|
|
cur.execute(
|
|
f"SELECT payload, "
|
|
f"1 - VEC_Distance_Cosine(embedding, VEC_FromText(%s)) "
|
|
f"FROM `{self._db}`.`{self._table}` "
|
|
f"ORDER BY VEC_Distance_Cosine(embedding, VEC_FromText(%s)) "
|
|
f"LIMIT %s", (qtext, qtext, limit))
|
|
rows = cur.fetchall()
|
|
return [{"score": float(score),
|
|
"payload": _snippet_payload(json.loads(payload))}
|
|
for payload, score in rows]
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────────
|
|
# Registry
|
|
# ──────────────────────────────────────────────────────────────────
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class TargetSpec:
|
|
key: str
|
|
label: str # UI display label
|
|
modules: tuple[str, ...] # probe requirements
|
|
service: str | None # live service required: "qdrant" | "postgres"
|
|
factory: Callable[..., object]
|
|
|
|
|
|
def _qdrant_factory(*, host: str, port: int, collection: str, dim: int,
|
|
log: Callable[[str], None], data_dir: Path):
|
|
from .qdrant_store import QdrantStore
|
|
return QdrantStore(host=host, port=port, collection=collection,
|
|
dim=dim, log=log)
|
|
|
|
|
|
def _weaviate_factory(*, host: str, port: int, collection: str, dim: int,
|
|
log: Callable[[str], None], data_dir: Path):
|
|
return WeaviateStore(host=host, port=port, collection=collection,
|
|
dim=dim, log=log)
|
|
|
|
|
|
def _embedded_factory(store_cls):
|
|
def factory(*, collection: str, dim: int,
|
|
log: Callable[[str], None], data_dir: Path, **_unused):
|
|
return store_cls(data_dir=data_dir, collection=collection,
|
|
dim=dim, log=log)
|
|
return factory
|
|
|
|
|
|
TARGETS: dict[str, TargetSpec] = {
|
|
spec.key: spec for spec in (
|
|
TargetSpec("qdrant", "Qdrant (service)", ("qdrant_client", "fastembed"),
|
|
"qdrant", _qdrant_factory),
|
|
TargetSpec("chroma", "Chroma (embedded)", ("chromadb",), None,
|
|
_embedded_factory(ChromaStore)),
|
|
TargetSpec("lancedb", "LanceDB (embedded)", ("lancedb",), None,
|
|
_embedded_factory(LanceStore)),
|
|
TargetSpec("faiss", "FAISS (file)", ("faiss",), None,
|
|
_embedded_factory(FaissStore)),
|
|
TargetSpec("milvus", "Milvus Lite (embedded)", ("pymilvus",), None,
|
|
_embedded_factory(MilvusStore)),
|
|
TargetSpec("weaviate", "Weaviate (service)", ("weaviate",),
|
|
"weaviate", _weaviate_factory),
|
|
TargetSpec("pgvector", "pgvector (Postgres)", ("psycopg", "pgvector"),
|
|
"postgres", _embedded_factory(PgVectorStore)),
|
|
TargetSpec("duckdb", "DuckDB (embedded)", ("duckdb",), None,
|
|
_embedded_factory(DuckStore)),
|
|
TargetSpec("sqlitevec", "sqlite-vec (embedded)", ("sqlite_vec",), None,
|
|
_embedded_factory(SqliteVecStore)),
|
|
TargetSpec("mariadb", "MariaDB (service)", ("pymysql",),
|
|
"mariadb", _embedded_factory(MariaDbStore)),
|
|
)
|
|
}
|
|
|
|
|
|
def create_store(target: str, *, collection: str, dim: int,
|
|
host: str = "localhost", port: int = 6333,
|
|
data_dir: Path | None = None,
|
|
log: Callable[[str], None] = lambda _msg: None):
|
|
"""Build the store for *target* — the single dispatch point every
|
|
caller (pipeline, retrieval, tests) shares."""
|
|
spec = TARGETS.get(target)
|
|
if spec is None:
|
|
known = ", ".join(sorted(TARGETS))
|
|
raise VectorStoreError(f"unknown vector target '{target}' — known: {known}")
|
|
return spec.factory(host=host, port=port, collection=collection,
|
|
dim=dim, log=log, data_dir=data_dir)
|