feat: add pluggable vector backends
This commit is contained in:
parent
9b7cfc9940
commit
6aa8e93bc9
22
README.md
22
README.md
|
|
@ -69,6 +69,28 @@ python -m venv .venv && source .venv/bin/activate
|
|||
pip install mempalace
|
||||
```
|
||||
|
||||
## Storage backends
|
||||
|
||||
ChromaDB is the default. For the pluggable-backend preview, MemPalace also
|
||||
ships `sqlite_exact` for local exact-vector correctness checks and `qdrant`
|
||||
for an opt-in Qdrant service backend.
|
||||
|
||||
```bash
|
||||
# local no-service backend
|
||||
mempalace mine ~/projects/myapp --backend sqlite_exact
|
||||
|
||||
# Qdrant backend, defaulting to http://localhost:6333
|
||||
MEMPALACE_QDRANT_URL=http://localhost:6333 \
|
||||
mempalace mine ~/projects/myapp --backend qdrant
|
||||
```
|
||||
|
||||
Qdrant can also be configured with `MEMPALACE_QDRANT_API_KEY`,
|
||||
`MEMPALACE_QDRANT_NAMESPACE`, and `MEMPALACE_QDRANT_TIMEOUT`.
|
||||
When `MEMPALACE_QDRANT_URL` points anywhere other than your own local or
|
||||
trusted self-hosted service, MemPalace will send and store verbatim drawer
|
||||
text and metadata there. That is an explicit opt-in backend choice, never
|
||||
the default.
|
||||
|
||||
## Quickstart
|
||||
|
||||
```bash
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ Public surface:
|
|||
from .base import (
|
||||
BackendClosedError,
|
||||
BackendError,
|
||||
BackendMismatchError,
|
||||
BaseBackend,
|
||||
BaseCollection,
|
||||
CollectionNotInitializedError,
|
||||
|
|
@ -24,14 +25,21 @@ from .base import (
|
|||
EmbedderIdentityMismatchError,
|
||||
GetResult,
|
||||
HealthStatus,
|
||||
LexicalHit,
|
||||
LexicalResult,
|
||||
PalaceNotFoundError,
|
||||
PalaceRef,
|
||||
QueryResult,
|
||||
UnsupportedCapabilityError,
|
||||
UnsupportedFilterError,
|
||||
)
|
||||
from .chroma import ChromaBackend, ChromaCollection
|
||||
from .qdrant import QdrantBackend, QdrantCollection
|
||||
from .sqlite_exact import SQLiteExactBackend, SQLiteExactCollection
|
||||
from .registry import (
|
||||
available_backends,
|
||||
detect_backend_for_path,
|
||||
detect_backends_for_path,
|
||||
get_backend,
|
||||
get_backend_class,
|
||||
register,
|
||||
|
|
@ -43,6 +51,7 @@ from .registry import (
|
|||
__all__ = [
|
||||
"BackendClosedError",
|
||||
"BackendError",
|
||||
"BackendMismatchError",
|
||||
"BaseBackend",
|
||||
"BaseCollection",
|
||||
"ChromaBackend",
|
||||
|
|
@ -52,11 +61,20 @@ __all__ = [
|
|||
"EmbedderIdentityMismatchError",
|
||||
"GetResult",
|
||||
"HealthStatus",
|
||||
"LexicalHit",
|
||||
"LexicalResult",
|
||||
"PalaceNotFoundError",
|
||||
"PalaceRef",
|
||||
"QdrantBackend",
|
||||
"QdrantCollection",
|
||||
"QueryResult",
|
||||
"SQLiteExactBackend",
|
||||
"SQLiteExactCollection",
|
||||
"UnsupportedCapabilityError",
|
||||
"UnsupportedFilterError",
|
||||
"available_backends",
|
||||
"detect_backend_for_path",
|
||||
"detect_backends_for_path",
|
||||
"get_backend",
|
||||
"get_backend_class",
|
||||
"register",
|
||||
|
|
|
|||
|
|
@ -58,6 +58,14 @@ class UnsupportedFilterError(BackendError):
|
|||
"""
|
||||
|
||||
|
||||
class UnsupportedCapabilityError(BackendError):
|
||||
"""Raised when a backend does not implement an optional capability."""
|
||||
|
||||
|
||||
class BackendMismatchError(BackendError):
|
||||
"""Raised when a selected backend does not match existing palace artifacts."""
|
||||
|
||||
|
||||
class DimensionMismatchError(BackendError):
|
||||
"""Raised when the embedding dimension on write does not match the collection."""
|
||||
|
||||
|
|
@ -177,6 +185,23 @@ class GetResult(_DictCompatMixin):
|
|||
return cls(ids=[], documents=[], metadatas=[], embeddings=None)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LexicalHit:
|
||||
"""One hit from backend lexical candidate search."""
|
||||
|
||||
id: str
|
||||
document: str
|
||||
metadata: dict
|
||||
score: float
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LexicalResult:
|
||||
"""Typed return from ``BaseCollection.lexical_search``."""
|
||||
|
||||
hits: list[LexicalHit]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Collection contract
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -253,6 +278,15 @@ class BaseCollection(ABC):
|
|||
def health(self) -> HealthStatus:
|
||||
return HealthStatus.healthy()
|
||||
|
||||
def lexical_search(
|
||||
self,
|
||||
*,
|
||||
query: str,
|
||||
n_results: int = 10,
|
||||
where: Optional[dict] = None,
|
||||
) -> LexicalResult:
|
||||
raise UnsupportedCapabilityError("backend does not support lexical_search")
|
||||
|
||||
def update(
|
||||
self,
|
||||
*,
|
||||
|
|
|
|||
|
|
@ -4,8 +4,10 @@ import contextlib
|
|||
import datetime as _dt
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import pickle
|
||||
import re
|
||||
import sqlite3
|
||||
from numbers import Integral
|
||||
from pathlib import Path
|
||||
|
|
@ -20,6 +22,8 @@ from .base import (
|
|||
CollectionNotInitializedError,
|
||||
GetResult,
|
||||
HealthStatus,
|
||||
LexicalHit,
|
||||
LexicalResult,
|
||||
PalaceNotFoundError,
|
||||
PalaceRef,
|
||||
QueryResult,
|
||||
|
|
@ -33,6 +37,7 @@ logger = logging.getLogger(__name__)
|
|||
_REQUIRED_OPERATORS = frozenset({"$eq", "$ne", "$in", "$nin", "$and", "$or", "$contains"})
|
||||
_OPTIONAL_OPERATORS = frozenset({"$gt", "$gte", "$lt", "$lte"})
|
||||
_SUPPORTED_OPERATORS = _REQUIRED_OPERATORS | _OPTIONAL_OPERATORS
|
||||
_TOKEN_RE = re.compile(r"\w{2,}", re.UNICODE)
|
||||
|
||||
# A healthy HNSW payload should keep link_lists.bin proportional to
|
||||
# data_level0.bin. When link_lists.bin grows orders of magnitude larger than
|
||||
|
|
@ -159,6 +164,127 @@ def _validate_where(where: Optional[dict]) -> None:
|
|||
stack.extend(x for x in v if isinstance(x, dict))
|
||||
|
||||
|
||||
def _tokenize(text: str) -> list[str]:
|
||||
if not text:
|
||||
return []
|
||||
return _TOKEN_RE.findall(text.lower())
|
||||
|
||||
|
||||
def _bm25_scores(
|
||||
query: str,
|
||||
documents: list[str],
|
||||
k1: float = 1.5,
|
||||
b: float = 0.75,
|
||||
) -> list[float]:
|
||||
query_terms = set(_tokenize(query))
|
||||
n_docs = len(documents)
|
||||
if not query_terms or n_docs == 0:
|
||||
return [0.0] * n_docs
|
||||
|
||||
tokenized = [_tokenize(doc) for doc in documents]
|
||||
doc_lens = [len(toks) for toks in tokenized]
|
||||
if not any(doc_lens):
|
||||
return [0.0] * n_docs
|
||||
avgdl = sum(doc_lens) / n_docs or 1.0
|
||||
|
||||
df = {term: 0 for term in query_terms}
|
||||
for toks in tokenized:
|
||||
for term in set(toks) & query_terms:
|
||||
df[term] += 1
|
||||
|
||||
idf = {
|
||||
term: math.log((n_docs - df[term] + 0.5) / (df[term] + 0.5) + 1.0) for term in query_terms
|
||||
}
|
||||
|
||||
scores = []
|
||||
for toks, dl in zip(tokenized, doc_lens):
|
||||
if dl == 0:
|
||||
scores.append(0.0)
|
||||
continue
|
||||
tf: dict[str, int] = {}
|
||||
for token in toks:
|
||||
if token in query_terms:
|
||||
tf[token] = tf.get(token, 0) + 1
|
||||
score = 0.0
|
||||
for term, freq in tf.items():
|
||||
num = freq * (k1 + 1)
|
||||
den = freq + k1 * (1 - b + b * dl / avgdl)
|
||||
score += idf[term] * num / den
|
||||
scores.append(score)
|
||||
return scores
|
||||
|
||||
|
||||
def _coerce_metadata_value(value: Any) -> Any:
|
||||
if isinstance(value, bool):
|
||||
return int(value)
|
||||
return value
|
||||
|
||||
|
||||
def _compare_metadata(actual: Any, op: str, expected: Any) -> bool:
|
||||
actual = _coerce_metadata_value(actual)
|
||||
expected = _coerce_metadata_value(expected)
|
||||
if op == "$eq":
|
||||
return actual == expected
|
||||
if op == "$ne":
|
||||
return actual != expected
|
||||
if op == "$in":
|
||||
return actual in (expected or [])
|
||||
if op == "$nin":
|
||||
return actual not in (expected or [])
|
||||
if op == "$contains":
|
||||
return str(expected) in str(actual or "")
|
||||
try:
|
||||
if op == "$gt":
|
||||
return actual > expected
|
||||
if op == "$gte":
|
||||
return actual >= expected
|
||||
if op == "$lt":
|
||||
return actual < expected
|
||||
if op == "$lte":
|
||||
return actual <= expected
|
||||
except TypeError:
|
||||
return False
|
||||
raise UnsupportedFilterError(f"operator {op!r} not supported by chroma backend")
|
||||
|
||||
|
||||
def _matches_where(meta: dict, where: Optional[dict]) -> bool:
|
||||
if not where:
|
||||
return True
|
||||
if not isinstance(where, dict):
|
||||
return False
|
||||
for key, expected in where.items():
|
||||
if key == "$and":
|
||||
if not all(_matches_where(meta, clause) for clause in expected or []):
|
||||
return False
|
||||
continue
|
||||
if key == "$or":
|
||||
if not any(_matches_where(meta, clause) for clause in expected or []):
|
||||
return False
|
||||
continue
|
||||
if key.startswith("$"):
|
||||
raise UnsupportedFilterError(f"operator {key!r} not supported by chroma backend")
|
||||
actual = meta.get(key)
|
||||
if isinstance(expected, dict):
|
||||
for op, operand in expected.items():
|
||||
if not _compare_metadata(actual, op, operand):
|
||||
return False
|
||||
elif actual != expected:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _metadata_cell_value(sval, ival, fval, bval):
|
||||
if sval is not None:
|
||||
return sval
|
||||
if ival is not None:
|
||||
return ival
|
||||
if fval is not None:
|
||||
return fval
|
||||
if bval is not None:
|
||||
return bool(bval)
|
||||
return None
|
||||
|
||||
|
||||
def _segment_appears_healthy(seg_dir: str) -> bool:
|
||||
"""Return True if a chromadb HNSW segment dir looks intact.
|
||||
|
||||
|
|
@ -1226,6 +1352,230 @@ class ChromaCollection(BaseCollection):
|
|||
def count(self):
|
||||
return self._collection.count()
|
||||
|
||||
def lexical_search(
|
||||
self,
|
||||
*,
|
||||
query: str,
|
||||
n_results: int = 10,
|
||||
where: Optional[dict] = None,
|
||||
) -> LexicalResult:
|
||||
"""Return lexical BM25 candidates for this collection.
|
||||
|
||||
This is the normal healthy-Chroma implementation behind the optional
|
||||
backend capability. The HNSW-disabled fallback in ``searcher.py`` still
|
||||
reads ``chroma.sqlite3`` directly and remains Chroma-only.
|
||||
"""
|
||||
_validate_where(where)
|
||||
sqlite_hits = self._lexical_search_via_sqlite(query=query, n_results=n_results, where=where)
|
||||
if sqlite_hits is not None:
|
||||
return LexicalResult(hits=sqlite_hits)
|
||||
|
||||
# Directly-constructed ChromaCollection test doubles may not carry a
|
||||
# palace path. Keep lexical_search usable in that shape, but normal
|
||||
# MemPalace paths above use Chroma's FTS table instead of scanning every
|
||||
# drawer through the Python client.
|
||||
total = self.count()
|
||||
docs: list[str] = []
|
||||
metas: list[dict] = []
|
||||
ids: list[str] = []
|
||||
offset = 0
|
||||
batch_size = 1000
|
||||
while offset < total:
|
||||
kwargs: dict[str, Any] = {
|
||||
"include": ["documents", "metadatas"],
|
||||
"limit": batch_size,
|
||||
"offset": offset,
|
||||
}
|
||||
if where:
|
||||
kwargs["where"] = where
|
||||
batch = self.get(**kwargs)
|
||||
if not batch.ids:
|
||||
break
|
||||
ids.extend(batch.ids)
|
||||
docs.extend(doc or "" for doc in batch.documents)
|
||||
metas.extend(meta or {} for meta in batch.metadatas)
|
||||
offset += len(batch.ids)
|
||||
|
||||
scores = _bm25_scores(query, docs)
|
||||
hits = [
|
||||
LexicalHit(id=doc_id, document=doc, metadata=meta, score=float(score))
|
||||
for doc_id, doc, meta, score in zip(ids, docs, metas, scores)
|
||||
if score > 0
|
||||
]
|
||||
hits.sort(key=lambda hit: hit.score, reverse=True)
|
||||
return LexicalResult(hits=hits[:n_results])
|
||||
|
||||
def _collection_name(self) -> Optional[str]:
|
||||
name = getattr(self._collection, "name", None)
|
||||
if callable(name):
|
||||
try:
|
||||
name = name()
|
||||
except TypeError:
|
||||
name = None
|
||||
return str(name) if name else None
|
||||
|
||||
def _lexical_search_via_sqlite(
|
||||
self,
|
||||
*,
|
||||
query: str,
|
||||
n_results: int,
|
||||
where: Optional[dict],
|
||||
max_candidates: int = 500,
|
||||
) -> Optional[list[LexicalHit]]:
|
||||
if not self._palace_path:
|
||||
return None
|
||||
db_path = os.path.join(self._palace_path, "chroma.sqlite3")
|
||||
if not os.path.isfile(db_path):
|
||||
return []
|
||||
collection_name = self._collection_name()
|
||||
if not collection_name:
|
||||
return []
|
||||
|
||||
tokens = [t for t in _tokenize(query) if len(t) >= 3]
|
||||
use_recency_fallback = not tokens
|
||||
candidate_ids: list[int] = []
|
||||
try:
|
||||
conn = sqlite3.connect(f"file:{db_path}?mode=ro", uri=True)
|
||||
conn.row_factory = sqlite3.Row
|
||||
except sqlite3.Error:
|
||||
logger.debug("Chroma lexical sqlite open failed", exc_info=True)
|
||||
return []
|
||||
|
||||
try:
|
||||
if tokens:
|
||||
fts_query = " OR ".join(tokens)
|
||||
# If a metadata filter is present, do not cap before filtering:
|
||||
# otherwise a common term can fill the window with wrong-scope
|
||||
# rows and hide valid scoped hits later in the FTS result set.
|
||||
limit_sql = "" if where else "LIMIT ?"
|
||||
params = [fts_query, collection_name]
|
||||
if not where:
|
||||
params.append(max(max_candidates, n_results))
|
||||
try:
|
||||
rows = conn.execute(
|
||||
f"""
|
||||
SELECT embedding_fulltext_search.rowid
|
||||
FROM embedding_fulltext_search
|
||||
JOIN embeddings e ON e.id = embedding_fulltext_search.rowid
|
||||
JOIN segments s ON e.segment_id = s.id
|
||||
JOIN collections c ON s.collection = c.id
|
||||
WHERE embedding_fulltext_search MATCH ?
|
||||
AND c.name = ?
|
||||
{limit_sql}
|
||||
""",
|
||||
params,
|
||||
).fetchall()
|
||||
candidate_ids = [int(row[0]) for row in rows]
|
||||
except sqlite3.Error:
|
||||
logger.debug(
|
||||
"Chroma lexical FTS query failed; using recency fallback", exc_info=True
|
||||
)
|
||||
use_recency_fallback = True
|
||||
|
||||
if not candidate_ids and use_recency_fallback:
|
||||
order_expr = "e.created_at DESC"
|
||||
try:
|
||||
rows = conn.execute(
|
||||
f"""
|
||||
SELECT e.id
|
||||
FROM embeddings e
|
||||
JOIN segments s ON e.segment_id = s.id
|
||||
JOIN collections c ON s.collection = c.id
|
||||
WHERE c.name = ?
|
||||
ORDER BY {order_expr}
|
||||
LIMIT ?
|
||||
""",
|
||||
(collection_name, max(max_candidates, n_results)),
|
||||
).fetchall()
|
||||
except sqlite3.Error:
|
||||
logger.debug(
|
||||
"Chroma lexical recency fallback failed; ordering by id", exc_info=True
|
||||
)
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT e.id
|
||||
FROM embeddings e
|
||||
JOIN segments s ON e.segment_id = s.id
|
||||
JOIN collections c ON s.collection = c.id
|
||||
WHERE c.name = ?
|
||||
ORDER BY e.id DESC
|
||||
LIMIT ?
|
||||
""",
|
||||
(collection_name, max(max_candidates, n_results)),
|
||||
).fetchall()
|
||||
candidate_ids = [int(row[0]) for row in rows]
|
||||
|
||||
if not candidate_ids:
|
||||
return []
|
||||
|
||||
meta_columns = {
|
||||
row["name"]
|
||||
for row in conn.execute("PRAGMA table_info(embedding_metadata)").fetchall()
|
||||
}
|
||||
value_columns = [
|
||||
col
|
||||
for col in ("string_value", "int_value", "float_value", "bool_value")
|
||||
if col in meta_columns
|
||||
]
|
||||
if not value_columns:
|
||||
return []
|
||||
meta_rows = []
|
||||
for start in range(0, len(candidate_ids), 900):
|
||||
chunk_ids = candidate_ids[start : start + 900]
|
||||
placeholders = ",".join("?" for _ in chunk_ids)
|
||||
meta_rows.extend(
|
||||
conn.execute(
|
||||
f"""
|
||||
SELECT id, key, {", ".join(value_columns)}
|
||||
FROM embedding_metadata
|
||||
WHERE id IN ({placeholders})
|
||||
""",
|
||||
chunk_ids,
|
||||
).fetchall()
|
||||
)
|
||||
except sqlite3.Error:
|
||||
logger.debug("Chroma lexical sqlite read failed", exc_info=True)
|
||||
return []
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
drawers: dict[int, dict] = {}
|
||||
for row in meta_rows:
|
||||
emb_id = int(row["id"])
|
||||
key = row["key"]
|
||||
values = {col: row[col] if col in row.keys() else None for col in value_columns}
|
||||
value = _metadata_cell_value(
|
||||
values.get("string_value"),
|
||||
values.get("int_value"),
|
||||
values.get("float_value"),
|
||||
values.get("bool_value"),
|
||||
)
|
||||
drawer = drawers.setdefault(emb_id, {"metadata": {}, "document": ""})
|
||||
if key == "chroma:document":
|
||||
drawer["document"] = str(value or "")
|
||||
else:
|
||||
drawer["metadata"][key] = value
|
||||
|
||||
ordered = []
|
||||
for emb_id in candidate_ids:
|
||||
drawer = drawers.get(emb_id)
|
||||
if drawer is None:
|
||||
continue
|
||||
meta = drawer["metadata"]
|
||||
if not _matches_where(meta, where):
|
||||
continue
|
||||
ordered.append((emb_id, drawer["document"], meta))
|
||||
|
||||
docs = [doc for _, doc, _ in ordered]
|
||||
scores = _bm25_scores(query, docs)
|
||||
hits = [
|
||||
LexicalHit(id=str(emb_id), document=doc, metadata=meta, score=float(score))
|
||||
for (emb_id, doc, meta), score in zip(ordered, scores)
|
||||
if score > 0
|
||||
]
|
||||
hits.sort(key=lambda hit: hit.score, reverse=True)
|
||||
return hits[:n_results]
|
||||
|
||||
@property
|
||||
def metadata(self) -> dict:
|
||||
"""Pass-through to the underlying ChromaDB collection's metadata.
|
||||
|
|
@ -1264,6 +1614,7 @@ class ChromaBackend(BaseBackend):
|
|||
"supports_embeddings_out",
|
||||
"supports_metadata_filters",
|
||||
"supports_contains_fast",
|
||||
"supports_lexical_search",
|
||||
"local_mode",
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,115 @@
|
|||
"""Core-side embedding adapter for explicit-vector backends."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from .base import BaseCollection
|
||||
|
||||
|
||||
def _embed_texts(texts: list[str]) -> list[list[float]]:
|
||||
"""Embed ``texts`` with the configured local embedding function."""
|
||||
if not texts:
|
||||
return []
|
||||
from ..embedding import get_embedding_function
|
||||
|
||||
ef = get_embedding_function()
|
||||
vectors = ef(input=texts)
|
||||
return [list(v) for v in vectors]
|
||||
|
||||
|
||||
class EmbeddingCollection(BaseCollection):
|
||||
"""Wrap a collection that requires explicit vectors.
|
||||
|
||||
Backends opt in with the ``requires_explicit_embeddings`` capability.
|
||||
Core callers can keep using ``documents=`` and ``query_texts=``; this
|
||||
wrapper computes vectors locally before delegating to the backend.
|
||||
"""
|
||||
|
||||
def __init__(self, inner: BaseCollection):
|
||||
self._inner = inner
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._inner, name)
|
||||
|
||||
def add(self, *, documents, ids, metadatas=None, embeddings=None):
|
||||
if embeddings is None:
|
||||
embeddings = _embed_texts(list(documents))
|
||||
return self._inner.add(
|
||||
documents=documents,
|
||||
ids=ids,
|
||||
metadatas=metadatas,
|
||||
embeddings=embeddings,
|
||||
)
|
||||
|
||||
def upsert(self, *, documents, ids, metadatas=None, embeddings=None):
|
||||
if embeddings is None:
|
||||
embeddings = _embed_texts(list(documents))
|
||||
return self._inner.upsert(
|
||||
documents=documents,
|
||||
ids=ids,
|
||||
metadatas=metadatas,
|
||||
embeddings=embeddings,
|
||||
)
|
||||
|
||||
def query(
|
||||
self,
|
||||
*,
|
||||
query_texts: Optional[list[str]] = None,
|
||||
query_embeddings: Optional[list[list[float]]] = None,
|
||||
n_results: int = 10,
|
||||
where: Optional[dict] = None,
|
||||
where_document: Optional[dict] = None,
|
||||
include: Optional[list[str]] = None,
|
||||
):
|
||||
if query_texts is not None and query_embeddings is None:
|
||||
query_embeddings = _embed_texts(list(query_texts))
|
||||
query_texts = None
|
||||
return self._inner.query(
|
||||
query_texts=query_texts,
|
||||
query_embeddings=query_embeddings,
|
||||
n_results=n_results,
|
||||
where=where,
|
||||
where_document=where_document,
|
||||
include=include,
|
||||
)
|
||||
|
||||
def get(
|
||||
self, *, ids=None, where=None, where_document=None, limit=None, offset=None, include=None
|
||||
):
|
||||
return self._inner.get(
|
||||
ids=ids,
|
||||
where=where,
|
||||
where_document=where_document,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
include=include,
|
||||
)
|
||||
|
||||
def delete(self, *, ids=None, where=None):
|
||||
return self._inner.delete(ids=ids, where=where)
|
||||
|
||||
def count(self) -> int:
|
||||
return self._inner.count()
|
||||
|
||||
def estimated_count(self) -> int:
|
||||
return self._inner.estimated_count()
|
||||
|
||||
def close(self) -> None:
|
||||
return self._inner.close()
|
||||
|
||||
def health(self):
|
||||
return self._inner.health()
|
||||
|
||||
def lexical_search(self, *, query: str, n_results: int = 10, where: Optional[dict] = None):
|
||||
return self._inner.lexical_search(query=query, n_results=n_results, where=where)
|
||||
|
||||
def update(self, *, ids, documents=None, metadatas=None, embeddings=None):
|
||||
if documents is not None and embeddings is None:
|
||||
embeddings = _embed_texts(list(documents))
|
||||
return self._inner.update(
|
||||
ids=ids,
|
||||
documents=documents,
|
||||
metadatas=metadatas,
|
||||
embeddings=embeddings,
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
|
|
@ -125,6 +125,38 @@ def get_backend(name: str) -> BaseBackend:
|
|||
return inst
|
||||
|
||||
|
||||
def detect_backends_for_path(path: str) -> list[str]:
|
||||
"""Return all registered backend names whose artifacts are present at ``path``.
|
||||
|
||||
Detection is a migration/protection aid for local palaces. Backends are
|
||||
checked in registry-name order so callers get deterministic diagnostics if
|
||||
a broken directory contains artifacts from more than one backend.
|
||||
"""
|
||||
_discover_entry_points()
|
||||
detected = []
|
||||
for name in sorted(_registry):
|
||||
cls = _registry[name]
|
||||
try:
|
||||
if cls.detect(path):
|
||||
detected.append(name)
|
||||
except Exception:
|
||||
logger.exception("detect() raised on backend %r", name)
|
||||
return detected
|
||||
|
||||
|
||||
def detect_backend_for_path(path: str) -> Optional[str]:
|
||||
"""Return the single detected backend at ``path``, or ``None``.
|
||||
|
||||
If multiple backend artifacts are present, the first name in registry order
|
||||
is returned for backward compatibility. Callers that enforce mismatch
|
||||
protection should use :func:`detect_backends_for_path`.
|
||||
"""
|
||||
detected = detect_backends_for_path(path)
|
||||
if detected:
|
||||
return detected[0]
|
||||
return None
|
||||
|
||||
|
||||
def reset_backends() -> None:
|
||||
"""Close and drop all cached backend instances (primarily for tests)."""
|
||||
with _lock:
|
||||
|
|
@ -161,14 +193,9 @@ def resolve_backend_for_palace(
|
|||
return candidate
|
||||
|
||||
_discover_entry_points()
|
||||
if palace_path:
|
||||
for name, cls in _registry.items():
|
||||
try:
|
||||
if cls.detect(palace_path):
|
||||
return name
|
||||
except Exception:
|
||||
logger.exception("detect() raised on backend %r", name)
|
||||
continue
|
||||
detected = detect_backend_for_path(palace_path) if palace_path else None
|
||||
if detected:
|
||||
return detected
|
||||
return default
|
||||
|
||||
|
||||
|
|
@ -180,10 +207,16 @@ def resolve_backend_for_palace(
|
|||
def _register_builtins() -> None:
|
||||
"""Register chroma as the in-tree default."""
|
||||
from .chroma import ChromaBackend
|
||||
from .qdrant import QdrantBackend
|
||||
from .sqlite_exact import SQLiteExactBackend
|
||||
|
||||
# Use setdefault semantics so a caller that pre-registered for tests wins.
|
||||
if "chroma" not in _registry:
|
||||
_registry["chroma"] = ChromaBackend
|
||||
if "qdrant" not in _registry:
|
||||
_registry["qdrant"] = QdrantBackend
|
||||
if "sqlite_exact" not in _registry:
|
||||
_registry["sqlite_exact"] = SQLiteExactBackend
|
||||
|
||||
|
||||
_register_builtins()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,940 @@
|
|||
"""SQLite exact-vector backend for MemPalace.
|
||||
|
||||
This backend is intentionally simple and local-first. It is a correctness
|
||||
backend, not a high-throughput ANN backend: vectors are stored as float32
|
||||
blobs and query uses exact cosine distance over the matching collection.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import sqlite3
|
||||
import threading
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .base import (
|
||||
BackendClosedError,
|
||||
BaseBackend,
|
||||
BaseCollection,
|
||||
CollectionNotInitializedError,
|
||||
DimensionMismatchError,
|
||||
GetResult,
|
||||
HealthStatus,
|
||||
LexicalHit,
|
||||
LexicalResult,
|
||||
PalaceNotFoundError,
|
||||
PalaceRef,
|
||||
QueryResult,
|
||||
UnsupportedFilterError,
|
||||
_IncludeSpec,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DB_FILENAME = "sqlite_exact.sqlite3"
|
||||
_TOKEN_RE = re.compile(r"\w{2,}", re.UNICODE)
|
||||
_SUPPORTED_OPERATORS = frozenset(
|
||||
{"$eq", "$ne", "$in", "$nin", "$and", "$or", "$contains", "$gt", "$gte", "$lt", "$lte"}
|
||||
)
|
||||
|
||||
|
||||
def _utcnow() -> str:
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
|
||||
|
||||
def _json_dumps(obj: Any) -> str:
|
||||
return json.dumps(obj or {}, ensure_ascii=False, separators=(",", ":"), sort_keys=True)
|
||||
|
||||
|
||||
def _json_loads(text: str | None) -> dict:
|
||||
if not text:
|
||||
return {}
|
||||
try:
|
||||
value = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
return {}
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _encode_vector(vector: list[float]) -> bytes:
|
||||
return _as_vector_array(vector).tobytes()
|
||||
|
||||
|
||||
def _as_vector_array(vector: list[float]) -> np.ndarray:
|
||||
arr = np.asarray(vector, dtype=np.float32)
|
||||
if arr.ndim != 1 or arr.size == 0:
|
||||
raise ValueError("embedding must be a non-empty 1D vector")
|
||||
return arr
|
||||
|
||||
|
||||
def _decode_vector(blob: bytes | None) -> list[float]:
|
||||
if not blob:
|
||||
return []
|
||||
return np.frombuffer(blob, dtype=np.float32).astype(float).tolist()
|
||||
|
||||
|
||||
def _decode_array(blob: bytes | None) -> Optional[np.ndarray]:
|
||||
if not blob:
|
||||
return None
|
||||
arr = np.frombuffer(blob, dtype=np.float32)
|
||||
if arr.size == 0:
|
||||
return None
|
||||
return arr
|
||||
|
||||
|
||||
def _tokenize(text: str) -> list[str]:
|
||||
if not text:
|
||||
return []
|
||||
return _TOKEN_RE.findall(text.lower())
|
||||
|
||||
|
||||
def _bm25_scores(query: str, documents: list[str], k1: float = 1.5, b: float = 0.75) -> list[float]:
|
||||
query_terms = set(_tokenize(query))
|
||||
n_docs = len(documents)
|
||||
if not query_terms or n_docs == 0:
|
||||
return [0.0] * n_docs
|
||||
|
||||
tokenized = [_tokenize(d) for d in documents]
|
||||
doc_lens = [len(toks) for toks in tokenized]
|
||||
if not any(doc_lens):
|
||||
return [0.0] * n_docs
|
||||
avgdl = sum(doc_lens) / n_docs or 1.0
|
||||
|
||||
df = {term: 0 for term in query_terms}
|
||||
for toks in tokenized:
|
||||
for term in set(toks) & query_terms:
|
||||
df[term] += 1
|
||||
|
||||
idf = {term: np.log((n_docs - df[term] + 0.5) / (df[term] + 0.5) + 1.0) for term in query_terms}
|
||||
|
||||
scores = []
|
||||
for toks, dl in zip(tokenized, doc_lens):
|
||||
if dl == 0:
|
||||
scores.append(0.0)
|
||||
continue
|
||||
tf: dict[str, int] = {}
|
||||
for token in toks:
|
||||
if token in query_terms:
|
||||
tf[token] = tf.get(token, 0) + 1
|
||||
score = 0.0
|
||||
for term, freq in tf.items():
|
||||
num = freq * (k1 + 1)
|
||||
den = freq + k1 * (1 - b + b * dl / avgdl)
|
||||
score += float(idf[term]) * num / den
|
||||
scores.append(score)
|
||||
return scores
|
||||
|
||||
|
||||
def _validate_where(where: Optional[dict]) -> None:
|
||||
if not where:
|
||||
return
|
||||
stack = [where]
|
||||
while stack:
|
||||
node = stack.pop()
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
for key, value in node.items():
|
||||
if key.startswith("$") and key not in _SUPPORTED_OPERATORS:
|
||||
raise UnsupportedFilterError(f"operator {key!r} not supported by sqlite_exact")
|
||||
if isinstance(value, dict):
|
||||
stack.append(value)
|
||||
elif isinstance(value, list):
|
||||
stack.extend(item for item in value if isinstance(item, dict))
|
||||
|
||||
|
||||
def _coerce_comparable(value: Any):
|
||||
if isinstance(value, bool):
|
||||
return int(value)
|
||||
return value
|
||||
|
||||
|
||||
def _compare(actual: Any, op: str, expected: Any) -> bool:
|
||||
actual = _coerce_comparable(actual)
|
||||
expected = _coerce_comparable(expected)
|
||||
if op == "$eq":
|
||||
return actual == expected
|
||||
if op == "$ne":
|
||||
return actual != expected
|
||||
if op == "$in":
|
||||
return actual in (expected or [])
|
||||
if op == "$nin":
|
||||
return actual not in (expected or [])
|
||||
if op == "$contains":
|
||||
return str(expected) in str(actual or "")
|
||||
try:
|
||||
if op == "$gt":
|
||||
return actual > expected
|
||||
if op == "$gte":
|
||||
return actual >= expected
|
||||
if op == "$lt":
|
||||
return actual < expected
|
||||
if op == "$lte":
|
||||
return actual <= expected
|
||||
except TypeError:
|
||||
return False
|
||||
raise UnsupportedFilterError(f"operator {op!r} not supported by sqlite_exact")
|
||||
|
||||
|
||||
def _matches_where(meta: dict, where: Optional[dict]) -> bool:
|
||||
if not where:
|
||||
return True
|
||||
if not isinstance(where, dict):
|
||||
return False
|
||||
for key, expected in where.items():
|
||||
if key == "$and":
|
||||
if not all(_matches_where(meta, clause) for clause in expected or []):
|
||||
return False
|
||||
continue
|
||||
if key == "$or":
|
||||
if not any(_matches_where(meta, clause) for clause in expected or []):
|
||||
return False
|
||||
continue
|
||||
if key.startswith("$"):
|
||||
raise UnsupportedFilterError(f"operator {key!r} not supported by sqlite_exact")
|
||||
actual = meta.get(key)
|
||||
if isinstance(expected, dict):
|
||||
for op, operand in expected.items():
|
||||
if not _compare(actual, op, operand):
|
||||
return False
|
||||
elif actual != expected:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _matches_where_document(document: str, where_document: Optional[dict]) -> bool:
|
||||
if not where_document:
|
||||
return True
|
||||
if not isinstance(where_document, dict):
|
||||
return False
|
||||
for key, value in where_document.items():
|
||||
if key == "$contains":
|
||||
if str(value) not in document:
|
||||
return False
|
||||
continue
|
||||
if key == "$and":
|
||||
if not all(_matches_where_document(document, clause) for clause in value or []):
|
||||
return False
|
||||
continue
|
||||
if key == "$or":
|
||||
if not any(_matches_where_document(document, clause) for clause in value or []):
|
||||
return False
|
||||
continue
|
||||
raise UnsupportedFilterError(f"where_document operator {key!r} not supported")
|
||||
return True
|
||||
|
||||
|
||||
def _validate_write_batch(
|
||||
*,
|
||||
documents: list[str],
|
||||
ids: list[str],
|
||||
metadatas: Optional[list[dict]],
|
||||
embeddings: Optional[list[list[float]]],
|
||||
) -> None:
|
||||
n = len(ids)
|
||||
if len(documents) != n:
|
||||
raise ValueError(f"documents length {len(documents)} does not match ids length {n}")
|
||||
if metadatas is not None and len(metadatas) != n:
|
||||
raise ValueError(f"metadatas length {len(metadatas)} does not match ids length {n}")
|
||||
if embeddings is not None and len(embeddings) != n:
|
||||
raise ValueError(f"embeddings length {len(embeddings)} does not match ids length {n}")
|
||||
|
||||
|
||||
class _SQLiteExactHandle:
|
||||
def __init__(self, conn: sqlite3.Connection, lock: threading.RLock):
|
||||
self.conn = conn
|
||||
self.lock = lock
|
||||
self.closed = False
|
||||
|
||||
|
||||
class SQLiteExactCollection(BaseCollection):
|
||||
def __init__(self, handle: _SQLiteExactHandle, collection_name: str):
|
||||
self._handle = handle
|
||||
self._collection_name = collection_name
|
||||
self._closed = False
|
||||
|
||||
def _ensure_open(self) -> None:
|
||||
if self._closed or self._handle.closed:
|
||||
raise BackendClosedError("SQLiteExactCollection has been closed")
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _cursor(self):
|
||||
with self._handle.lock:
|
||||
self._ensure_open()
|
||||
cur = self._handle.conn.cursor()
|
||||
try:
|
||||
yield cur
|
||||
except Exception:
|
||||
self._handle.conn.rollback()
|
||||
raise
|
||||
else:
|
||||
self._handle.conn.commit()
|
||||
finally:
|
||||
cur.close()
|
||||
|
||||
def _collection_id(self, cur) -> int:
|
||||
row = cur.execute(
|
||||
"SELECT id FROM collections WHERE name = ?",
|
||||
(self._collection_name,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise CollectionNotInitializedError(self._collection_name)
|
||||
return int(row[0])
|
||||
|
||||
def _collection_dimension(self, cur, collection_id: int) -> Optional[int]:
|
||||
row = cur.execute(
|
||||
"SELECT dimension FROM collections WHERE id = ?",
|
||||
(collection_id,),
|
||||
).fetchone()
|
||||
if row is None or row[0] is None:
|
||||
return None
|
||||
return int(row[0])
|
||||
|
||||
def _ensure_collection_dimension(self, cur, collection_id: int, dims: list[int]) -> None:
|
||||
distinct = {int(dim) for dim in dims}
|
||||
if not distinct:
|
||||
return
|
||||
if len(distinct) > 1:
|
||||
raise DimensionMismatchError(
|
||||
f"sqlite_exact collection {self._collection_name!r} cannot mix "
|
||||
f"embedding dimensions {sorted(distinct)}"
|
||||
)
|
||||
dim = distinct.pop()
|
||||
stored = self._collection_dimension(cur, collection_id)
|
||||
if stored is None:
|
||||
cur.execute(
|
||||
"UPDATE collections SET dimension = ? WHERE id = ?",
|
||||
(dim, collection_id),
|
||||
)
|
||||
elif stored != dim:
|
||||
raise DimensionMismatchError(
|
||||
f"sqlite_exact collection {self._collection_name!r} expects "
|
||||
f"embedding dimension {stored}, got {dim}"
|
||||
)
|
||||
|
||||
def _fts_available(self, cur) -> bool:
|
||||
row = cur.execute("SELECT value FROM meta WHERE key = 'fts5_available'").fetchone()
|
||||
return bool(row and row[0] == "1")
|
||||
|
||||
def _replace_fts(self, cur, collection_id: int, doc_id: str, document: str) -> None:
|
||||
if not self._fts_available(cur):
|
||||
return
|
||||
cur.execute(
|
||||
"DELETE FROM docs_fts WHERE collection_id = ? AND doc_id = ?",
|
||||
(collection_id, doc_id),
|
||||
)
|
||||
cur.execute(
|
||||
"INSERT INTO docs_fts(collection_id, doc_id, document) VALUES (?, ?, ?)",
|
||||
(collection_id, doc_id, document),
|
||||
)
|
||||
|
||||
def add(self, *, documents, ids, metadatas=None, embeddings=None):
|
||||
_validate_write_batch(
|
||||
documents=documents,
|
||||
ids=ids,
|
||||
metadatas=metadatas,
|
||||
embeddings=embeddings,
|
||||
)
|
||||
if embeddings is None:
|
||||
raise ValueError("sqlite_exact requires explicit embeddings")
|
||||
metadatas = metadatas or [{} for _ in ids]
|
||||
now = _utcnow()
|
||||
with self._cursor() as cur:
|
||||
collection_id = self._collection_id(cur)
|
||||
prepared = []
|
||||
for doc_id, doc, meta, emb in zip(ids, documents, metadatas, embeddings):
|
||||
arr = _as_vector_array(emb)
|
||||
prepared.append((doc_id, doc, meta, arr.tobytes(), int(arr.size)))
|
||||
self._ensure_collection_dimension(cur, collection_id, [item[4] for item in prepared])
|
||||
for doc_id, doc, meta, emb_blob, dim in prepared:
|
||||
cur.execute(
|
||||
"""
|
||||
INSERT INTO documents
|
||||
(collection_id, id, document, metadata_json, embedding, dim, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
collection_id,
|
||||
doc_id,
|
||||
doc,
|
||||
_json_dumps(meta),
|
||||
emb_blob,
|
||||
dim,
|
||||
now,
|
||||
now,
|
||||
),
|
||||
)
|
||||
self._replace_fts(cur, collection_id, doc_id, doc)
|
||||
|
||||
def upsert(self, *, documents, ids, metadatas=None, embeddings=None):
|
||||
_validate_write_batch(
|
||||
documents=documents,
|
||||
ids=ids,
|
||||
metadatas=metadatas,
|
||||
embeddings=embeddings,
|
||||
)
|
||||
if embeddings is None:
|
||||
raise ValueError("sqlite_exact requires explicit embeddings")
|
||||
metadatas = metadatas or [{} for _ in ids]
|
||||
now = _utcnow()
|
||||
with self._cursor() as cur:
|
||||
collection_id = self._collection_id(cur)
|
||||
prepared = []
|
||||
for doc_id, doc, meta, emb in zip(ids, documents, metadatas, embeddings):
|
||||
arr = _as_vector_array(emb)
|
||||
prepared.append((doc_id, doc, meta, arr.tobytes(), int(arr.size)))
|
||||
self._ensure_collection_dimension(cur, collection_id, [item[4] for item in prepared])
|
||||
for doc_id, doc, meta, emb_blob, dim in prepared:
|
||||
cur.execute(
|
||||
"""
|
||||
INSERT INTO documents
|
||||
(collection_id, id, document, metadata_json, embedding, dim, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(collection_id, id) DO UPDATE SET
|
||||
document = excluded.document,
|
||||
metadata_json = excluded.metadata_json,
|
||||
embedding = excluded.embedding,
|
||||
dim = excluded.dim,
|
||||
updated_at = excluded.updated_at
|
||||
""",
|
||||
(
|
||||
collection_id,
|
||||
doc_id,
|
||||
doc,
|
||||
_json_dumps(meta),
|
||||
emb_blob,
|
||||
dim,
|
||||
now,
|
||||
now,
|
||||
),
|
||||
)
|
||||
self._replace_fts(cur, collection_id, doc_id, doc)
|
||||
|
||||
def update(self, *, ids, documents=None, metadatas=None, embeddings=None):
|
||||
if documents is None and metadatas is None and embeddings is None:
|
||||
raise ValueError("update requires at least one of documents, metadatas, embeddings")
|
||||
n = len(ids)
|
||||
for label, value in (
|
||||
("documents", documents),
|
||||
("metadatas", metadatas),
|
||||
("embeddings", embeddings),
|
||||
):
|
||||
if value is not None and len(value) != n:
|
||||
raise ValueError(f"{label} length {len(value)} does not match ids length {n}")
|
||||
with self._cursor() as cur:
|
||||
collection_id = self._collection_id(cur)
|
||||
updates = []
|
||||
for idx, doc_id in enumerate(ids):
|
||||
row = cur.execute(
|
||||
"""
|
||||
SELECT document, metadata_json, embedding, dim
|
||||
FROM documents
|
||||
WHERE collection_id = ? AND id = ?
|
||||
""",
|
||||
(collection_id, doc_id),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
continue
|
||||
doc = documents[idx] if documents is not None else row[0]
|
||||
meta = _json_loads(row[1])
|
||||
if metadatas is not None:
|
||||
meta.update(metadatas[idx] or {})
|
||||
if embeddings is not None:
|
||||
arr = _as_vector_array(embeddings[idx])
|
||||
emb_blob = arr.tobytes()
|
||||
dim = int(arr.size)
|
||||
else:
|
||||
emb_blob = row[2]
|
||||
dim = row[3]
|
||||
updates.append((doc_id, doc, meta, emb_blob, dim))
|
||||
if embeddings is not None:
|
||||
self._ensure_collection_dimension(cur, collection_id, [item[4] for item in updates])
|
||||
for doc_id, doc, meta, emb_blob, dim in updates:
|
||||
cur.execute(
|
||||
"""
|
||||
UPDATE documents
|
||||
SET document = ?, metadata_json = ?, embedding = ?, dim = ?, updated_at = ?
|
||||
WHERE collection_id = ? AND id = ?
|
||||
""",
|
||||
(doc, _json_dumps(meta), emb_blob, dim, _utcnow(), collection_id, doc_id),
|
||||
)
|
||||
self._replace_fts(cur, collection_id, doc_id, doc)
|
||||
|
||||
def _rows(self, cur, *, where=None, where_document=None) -> list[dict]:
|
||||
_validate_where(where)
|
||||
_validate_where(where_document)
|
||||
collection_id = self._collection_id(cur)
|
||||
rows = cur.execute(
|
||||
"""
|
||||
SELECT id, document, metadata_json, embedding
|
||||
FROM documents
|
||||
WHERE collection_id = ?
|
||||
ORDER BY rowid
|
||||
""",
|
||||
(collection_id,),
|
||||
).fetchall()
|
||||
out = []
|
||||
for doc_id, doc, meta_json, emb_blob in rows:
|
||||
meta = _json_loads(meta_json)
|
||||
if not _matches_where(meta, where):
|
||||
continue
|
||||
if not _matches_where_document(doc or "", where_document):
|
||||
continue
|
||||
out.append(
|
||||
{
|
||||
"id": doc_id,
|
||||
"document": doc or "",
|
||||
"metadata": meta,
|
||||
"embedding": emb_blob,
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
def query(
|
||||
self,
|
||||
*,
|
||||
query_texts=None,
|
||||
query_embeddings=None,
|
||||
n_results=10,
|
||||
where=None,
|
||||
where_document=None,
|
||||
include=None,
|
||||
) -> QueryResult:
|
||||
if query_texts is not None:
|
||||
raise ValueError(
|
||||
"sqlite_exact requires query_embeddings; use palace.get_collection wrapper"
|
||||
)
|
||||
if query_embeddings is None:
|
||||
raise ValueError("query requires query_embeddings")
|
||||
if not query_embeddings:
|
||||
raise ValueError("query input must be a non-empty list")
|
||||
|
||||
spec = _IncludeSpec.resolve(include, default_distances=True)
|
||||
outer_ids: list[list[str]] = []
|
||||
outer_docs: list[list[str]] = []
|
||||
outer_metas: list[list[dict]] = []
|
||||
outer_dists: list[list[float]] = []
|
||||
outer_embeds: list[list[list[float]]] = []
|
||||
|
||||
with self._cursor() as cur:
|
||||
collection_id = self._collection_id(cur)
|
||||
expected_dim = self._collection_dimension(cur, collection_id)
|
||||
rows = self._rows(cur, where=where, where_document=where_document)
|
||||
row_vectors = [(row, _decode_array(row["embedding"])) for row in rows]
|
||||
|
||||
for query_vector in query_embeddings:
|
||||
q = _as_vector_array(query_vector)
|
||||
if expected_dim is not None and int(q.size) != expected_dim:
|
||||
raise DimensionMismatchError(
|
||||
f"sqlite_exact collection {self._collection_name!r} expects "
|
||||
f"embedding dimension {expected_dim}, got {int(q.size)}"
|
||||
)
|
||||
q_norm = float(np.linalg.norm(q))
|
||||
scored = []
|
||||
for row, vec in row_vectors:
|
||||
if vec is None or vec.size != q.size:
|
||||
continue
|
||||
denom = q_norm * float(np.linalg.norm(vec))
|
||||
cos = 0.0 if denom <= 0 else float(np.dot(q, vec) / denom)
|
||||
distance = 1.0 - max(-1.0, min(1.0, cos))
|
||||
scored.append((distance, row, vec))
|
||||
scored.sort(key=lambda item: item[0])
|
||||
top = scored[:n_results]
|
||||
|
||||
outer_ids.append([row["id"] for _, row, _ in top])
|
||||
outer_docs.append([row["document"] for _, row, _ in top] if spec.documents else [])
|
||||
outer_metas.append([row["metadata"] for _, row, _ in top] if spec.metadatas else [])
|
||||
outer_dists.append([float(dist) for dist, _, _ in top] if spec.distances else [])
|
||||
if spec.embeddings:
|
||||
outer_embeds.append([vec.astype(float).tolist() for _, _, vec in top])
|
||||
|
||||
return QueryResult(
|
||||
ids=outer_ids,
|
||||
documents=outer_docs,
|
||||
metadatas=outer_metas,
|
||||
distances=outer_dists,
|
||||
embeddings=outer_embeds if spec.embeddings else None,
|
||||
)
|
||||
|
||||
def get(
|
||||
self,
|
||||
*,
|
||||
ids=None,
|
||||
where=None,
|
||||
where_document=None,
|
||||
limit=None,
|
||||
offset=None,
|
||||
include=None,
|
||||
) -> GetResult:
|
||||
spec = _IncludeSpec.resolve(include, default_distances=False)
|
||||
with self._cursor() as cur:
|
||||
rows = self._rows(cur, where=where, where_document=where_document)
|
||||
if ids is not None:
|
||||
by_id = {row["id"]: row for row in rows}
|
||||
rows = [by_id[doc_id] for doc_id in ids if doc_id in by_id]
|
||||
if offset:
|
||||
rows = rows[offset:]
|
||||
if limit is not None:
|
||||
rows = rows[:limit]
|
||||
return GetResult(
|
||||
ids=[row["id"] for row in rows],
|
||||
documents=[row["document"] for row in rows] if spec.documents else [],
|
||||
metadatas=[row["metadata"] for row in rows] if spec.metadatas else [],
|
||||
embeddings=(
|
||||
[_decode_vector(row["embedding"]) for row in rows] if spec.embeddings else None
|
||||
),
|
||||
)
|
||||
|
||||
def delete(self, *, ids=None, where=None):
|
||||
with self._cursor() as cur:
|
||||
collection_id = self._collection_id(cur)
|
||||
if ids is None:
|
||||
rows = self._rows(cur, where=where)
|
||||
ids = [row["id"] for row in rows]
|
||||
for doc_id in ids or []:
|
||||
cur.execute(
|
||||
"DELETE FROM documents WHERE collection_id = ? AND id = ?",
|
||||
(collection_id, doc_id),
|
||||
)
|
||||
if self._fts_available(cur):
|
||||
cur.execute(
|
||||
"DELETE FROM docs_fts WHERE collection_id = ? AND doc_id = ?",
|
||||
(collection_id, doc_id),
|
||||
)
|
||||
|
||||
def count(self) -> int:
|
||||
with self._cursor() as cur:
|
||||
collection_id = self._collection_id(cur)
|
||||
row = cur.execute(
|
||||
"SELECT COUNT(*) FROM documents WHERE collection_id = ?",
|
||||
(collection_id,),
|
||||
).fetchone()
|
||||
return int(row[0]) if row else 0
|
||||
|
||||
def lexical_search(self, *, query: str, n_results: int = 10, where: Optional[dict] = None):
|
||||
_validate_where(where)
|
||||
with self._cursor() as cur:
|
||||
hits = self._lexical_search_fts(cur, query=query, n_results=n_results, where=where)
|
||||
if hits is not None:
|
||||
return LexicalResult(hits=hits)
|
||||
rows = self._rows(cur, where=where)
|
||||
scores = _bm25_scores(query, [row["document"] for row in rows])
|
||||
scored = [
|
||||
LexicalHit(
|
||||
id=row["id"],
|
||||
document=row["document"],
|
||||
metadata=row["metadata"],
|
||||
score=score,
|
||||
)
|
||||
for row, score in zip(rows, scores)
|
||||
if score > 0
|
||||
]
|
||||
scored.sort(key=lambda hit: hit.score, reverse=True)
|
||||
return LexicalResult(hits=scored[:n_results])
|
||||
|
||||
def _lexical_search_fts(self, cur, *, query: str, n_results: int, where: Optional[dict]):
|
||||
if not self._fts_available(cur):
|
||||
return None
|
||||
tokens = [t for t in _tokenize(query) if len(t) >= 2]
|
||||
if not tokens:
|
||||
return None
|
||||
fts_query = " OR ".join(tokens)
|
||||
collection_id = self._collection_id(cur)
|
||||
try:
|
||||
limit_sql = "" if where else "LIMIT ?"
|
||||
params = (fts_query, collection_id)
|
||||
if not where:
|
||||
params = (*params, max(n_results * 5, n_results))
|
||||
rows = cur.execute(
|
||||
f"""
|
||||
SELECT doc_id, bm25(docs_fts) AS rank
|
||||
FROM docs_fts
|
||||
WHERE docs_fts MATCH ? AND collection_id = ?
|
||||
ORDER BY rank
|
||||
{limit_sql}
|
||||
""",
|
||||
params,
|
||||
).fetchall()
|
||||
except sqlite3.Error:
|
||||
logger.debug("sqlite_exact FTS query failed; using Python lexical scan", exc_info=True)
|
||||
return None
|
||||
if not rows:
|
||||
return []
|
||||
ids = [row[0] for row in rows]
|
||||
docs = []
|
||||
for start in range(0, len(ids), 900):
|
||||
chunk_ids = ids[start : start + 900]
|
||||
placeholders = ",".join("?" for _ in chunk_ids)
|
||||
docs.extend(
|
||||
cur.execute(
|
||||
f"""
|
||||
SELECT id, document, metadata_json
|
||||
FROM documents
|
||||
WHERE collection_id = ? AND id IN ({placeholders})
|
||||
""",
|
||||
(collection_id, *chunk_ids),
|
||||
).fetchall()
|
||||
)
|
||||
by_id = {doc_id: (doc or "", _json_loads(meta_json)) for doc_id, doc, meta_json in docs}
|
||||
hits = []
|
||||
for doc_id, rank in rows:
|
||||
doc_meta = by_id.get(doc_id)
|
||||
if doc_meta is None:
|
||||
continue
|
||||
doc, meta = doc_meta
|
||||
if not _matches_where(meta, where):
|
||||
continue
|
||||
hits.append(
|
||||
LexicalHit(
|
||||
id=doc_id,
|
||||
document=doc,
|
||||
metadata=meta,
|
||||
score=-float(rank),
|
||||
)
|
||||
)
|
||||
if len(hits) >= n_results:
|
||||
break
|
||||
return hits
|
||||
|
||||
def close(self) -> None:
|
||||
self._closed = True
|
||||
|
||||
def health(self) -> HealthStatus:
|
||||
if self._closed or self._handle.closed:
|
||||
return HealthStatus.unhealthy("collection closed")
|
||||
return HealthStatus.healthy()
|
||||
|
||||
|
||||
class SQLiteExactBackend(BaseBackend):
|
||||
name = "sqlite_exact"
|
||||
capabilities = frozenset(
|
||||
{
|
||||
"requires_explicit_embeddings",
|
||||
"supports_embeddings_in",
|
||||
"supports_embeddings_passthrough",
|
||||
"supports_embeddings_out",
|
||||
"supports_metadata_filters",
|
||||
"supports_lexical_search",
|
||||
"local_mode",
|
||||
}
|
||||
)
|
||||
|
||||
def __init__(self):
|
||||
self._clients: dict[str, _SQLiteExactHandle] = {}
|
||||
self._clients_lock = threading.RLock()
|
||||
self._closed = False
|
||||
|
||||
@staticmethod
|
||||
def _db_path(palace_path: str) -> str:
|
||||
return os.path.join(palace_path, _DB_FILENAME)
|
||||
|
||||
def _connect(self, palace_path: str, create: bool):
|
||||
if self._closed:
|
||||
raise BackendClosedError("SQLiteExactBackend has been closed")
|
||||
db_path = self._db_path(palace_path)
|
||||
if not create and not os.path.isfile(db_path):
|
||||
raise PalaceNotFoundError(db_path)
|
||||
if create:
|
||||
os.makedirs(palace_path, exist_ok=True)
|
||||
try:
|
||||
os.chmod(palace_path, 0o700)
|
||||
except (OSError, NotImplementedError):
|
||||
pass
|
||||
with self._clients_lock:
|
||||
cached = self._clients.get(palace_path)
|
||||
if cached is not None and not cached.closed:
|
||||
return cached
|
||||
conn = sqlite3.connect(db_path, check_same_thread=False)
|
||||
conn.row_factory = sqlite3.Row
|
||||
lock = threading.RLock()
|
||||
handle = _SQLiteExactHandle(conn, lock)
|
||||
with handle.lock:
|
||||
self._init_schema(conn)
|
||||
with self._clients_lock:
|
||||
self._clients[palace_path] = handle
|
||||
return handle
|
||||
|
||||
def _init_schema(self, conn: sqlite3.Connection) -> None:
|
||||
conn.executescript(
|
||||
"""
|
||||
PRAGMA journal_mode=WAL;
|
||||
CREATE TABLE IF NOT EXISTS meta (
|
||||
key TEXT PRIMARY KEY,
|
||||
value TEXT NOT NULL
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS collections (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
dimension INTEGER,
|
||||
created_at TEXT NOT NULL
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS documents (
|
||||
collection_id INTEGER NOT NULL,
|
||||
id TEXT NOT NULL,
|
||||
document TEXT NOT NULL,
|
||||
metadata_json TEXT NOT NULL,
|
||||
embedding BLOB NOT NULL,
|
||||
dim INTEGER NOT NULL,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL,
|
||||
PRIMARY KEY (collection_id, id),
|
||||
FOREIGN KEY(collection_id) REFERENCES collections(id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_documents_collection
|
||||
ON documents(collection_id);
|
||||
"""
|
||||
)
|
||||
columns = {row[1] for row in conn.execute("PRAGMA table_info(collections)").fetchall()}
|
||||
if "dimension" not in columns:
|
||||
conn.execute("ALTER TABLE collections ADD COLUMN dimension INTEGER")
|
||||
try:
|
||||
conn.execute(
|
||||
"""
|
||||
CREATE VIRTUAL TABLE IF NOT EXISTS docs_fts
|
||||
USING fts5(collection_id UNINDEXED, doc_id UNINDEXED, document)
|
||||
"""
|
||||
)
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO meta(key, value)
|
||||
VALUES ('fts5_available', '1')
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value
|
||||
"""
|
||||
)
|
||||
except sqlite3.OperationalError:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO meta(key, value)
|
||||
VALUES ('fts5_available', '0')
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value
|
||||
"""
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
def get_collection(
|
||||
self,
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> SQLiteExactCollection:
|
||||
palace, collection_name, create = self._normalize_args(args, kwargs)
|
||||
palace_path = palace.local_path
|
||||
if palace_path is None:
|
||||
raise PalaceNotFoundError("SQLiteExactBackend requires PalaceRef.local_path")
|
||||
if not create and not os.path.isdir(palace_path):
|
||||
raise PalaceNotFoundError(palace_path)
|
||||
handle = self._connect(palace_path, create=create)
|
||||
with handle.lock:
|
||||
row = handle.conn.execute(
|
||||
"SELECT id FROM collections WHERE name = ?",
|
||||
(collection_name,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
if not create:
|
||||
raise CollectionNotInitializedError(palace_path)
|
||||
handle.conn.execute(
|
||||
"INSERT INTO collections(name, created_at) VALUES (?, ?)",
|
||||
(collection_name, _utcnow()),
|
||||
)
|
||||
handle.conn.commit()
|
||||
return SQLiteExactCollection(handle, collection_name)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_args(args, kwargs):
|
||||
if "palace" in kwargs:
|
||||
palace = kwargs.pop("palace")
|
||||
if not isinstance(palace, PalaceRef):
|
||||
raise TypeError("palace= must be a PalaceRef instance")
|
||||
collection_name = kwargs.pop("collection_name")
|
||||
create = bool(kwargs.pop("create", False))
|
||||
kwargs.pop("options", None)
|
||||
if args or kwargs:
|
||||
raise TypeError("unexpected arguments to get_collection")
|
||||
return palace, collection_name, create
|
||||
if args:
|
||||
palace_path = args[0]
|
||||
rest = list(args[1:])
|
||||
collection_name = kwargs.pop("collection_name", None) or (rest.pop(0) if rest else None)
|
||||
if collection_name is None:
|
||||
raise TypeError("collection_name is required")
|
||||
create = kwargs.pop("create", False)
|
||||
if rest:
|
||||
create = rest.pop(0)
|
||||
if rest or kwargs:
|
||||
raise TypeError("unexpected arguments to get_collection")
|
||||
return PalaceRef(id=palace_path, local_path=palace_path), collection_name, bool(create)
|
||||
if "palace_path" in kwargs:
|
||||
palace_path = kwargs.pop("palace_path")
|
||||
collection_name = kwargs.pop("collection_name")
|
||||
create = bool(kwargs.pop("create", False))
|
||||
if kwargs:
|
||||
raise TypeError("unexpected arguments to get_collection")
|
||||
return PalaceRef(id=palace_path, local_path=palace_path), collection_name, create
|
||||
raise TypeError("get_collection requires palace= or a positional palace_path")
|
||||
|
||||
def close_palace(self, palace: PalaceRef | str) -> None:
|
||||
path = palace.local_path if isinstance(palace, PalaceRef) else palace
|
||||
if path is None:
|
||||
return
|
||||
with self._clients_lock:
|
||||
cached = self._clients.pop(path, None)
|
||||
if cached is not None:
|
||||
with cached.lock:
|
||||
cached.closed = True
|
||||
cached.conn.close()
|
||||
|
||||
def close(self) -> None:
|
||||
with self._clients_lock:
|
||||
handles = list(self._clients.values())
|
||||
self._clients.clear()
|
||||
for handle in handles:
|
||||
with handle.lock:
|
||||
handle.closed = True
|
||||
handle.conn.close()
|
||||
self._closed = True
|
||||
|
||||
def health(self, palace: Optional[PalaceRef] = None) -> HealthStatus:
|
||||
if self._closed:
|
||||
return HealthStatus.unhealthy("backend closed")
|
||||
if palace and palace.local_path and not os.path.isfile(self._db_path(palace.local_path)):
|
||||
return HealthStatus.unhealthy("sqlite_exact database not found")
|
||||
return HealthStatus.healthy()
|
||||
|
||||
@classmethod
|
||||
def detect(cls, path: str) -> bool:
|
||||
return os.path.isfile(os.path.join(path, _DB_FILENAME))
|
||||
|
||||
def create_collection(self, palace_path: str, collection_name: str) -> SQLiteExactCollection:
|
||||
return self.get_collection(palace_path, collection_name, create=True)
|
||||
|
||||
def get_or_create_collection(self, palace_path: str, collection_name: str):
|
||||
return self.get_collection(palace_path, collection_name, create=True)
|
||||
|
||||
def delete_collection(self, palace_path: str, collection_name: str) -> None:
|
||||
handle = self._connect(palace_path, create=False)
|
||||
with handle.lock:
|
||||
row = handle.conn.execute(
|
||||
"SELECT id FROM collections WHERE name = ?",
|
||||
(collection_name,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise CollectionNotInitializedError(palace_path)
|
||||
collection_id = int(row[0])
|
||||
handle.conn.execute("DELETE FROM documents WHERE collection_id = ?", (collection_id,))
|
||||
try:
|
||||
handle.conn.execute(
|
||||
"DELETE FROM docs_fts WHERE collection_id = ?",
|
||||
(collection_id,),
|
||||
)
|
||||
except sqlite3.OperationalError:
|
||||
pass
|
||||
handle.conn.execute("DELETE FROM collections WHERE id = ?", (collection_id,))
|
||||
handle.conn.commit()
|
||||
|
||||
|
||||
__all__ = ["SQLiteExactBackend", "SQLiteExactCollection"]
|
||||
123
mempalace/cli.py
123
mempalace/cli.py
|
|
@ -51,6 +51,45 @@ _PASS_ZERO_PER_FILE_CAP = 100_000 # 100KB per file is generous for prose
|
|||
_PASS_ZERO_TOTAL_CAP = 5_000_000 # 5MB total ceiling — bounds memory
|
||||
_PASS_ZERO_LLM_PER_SAMPLE = 2_000 # for Tier 2 LLM call only
|
||||
_PASS_ZERO_LLM_MAX_SAMPLES = 20 # caps the LLM-tier sample count
|
||||
_EXPLICIT_BACKEND_ENV = "MEMPALACE_BACKEND_EXPLICIT"
|
||||
|
||||
|
||||
def _backend_arg(args):
|
||||
"""Return a CLI-selected backend from subcommand or global flags."""
|
||||
return getattr(args, "backend", None) or getattr(args, "global_backend", None)
|
||||
|
||||
|
||||
def _apply_backend_arg(args) -> None:
|
||||
backend = _backend_arg(args)
|
||||
if not backend:
|
||||
return
|
||||
backend = str(backend).strip().lower()
|
||||
from .backends import get_backend_class
|
||||
|
||||
get_backend_class(backend)
|
||||
os.environ[_EXPLICIT_BACKEND_ENV] = backend
|
||||
os.environ["MEMPALACE_BACKEND"] = backend
|
||||
|
||||
|
||||
def _selected_backend_for_palace(palace_path: str) -> str:
|
||||
from .palace import resolve_backend_name
|
||||
|
||||
return resolve_backend_name(palace_path, explicit=os.environ.get(_EXPLICIT_BACKEND_ENV))
|
||||
|
||||
|
||||
def _maintenance_requires_chroma(palace_path: str, command_name: str) -> bool:
|
||||
try:
|
||||
backend_name = _selected_backend_for_palace(palace_path)
|
||||
except Exception as exc: # noqa: BLE001 - user-facing guard before maintenance imports
|
||||
print(f"\n {command_name} cannot resolve the palace backend: {exc}", file=sys.stderr)
|
||||
return False
|
||||
if backend_name == "chroma":
|
||||
return True
|
||||
print(
|
||||
f"\n {command_name} is Chroma-only in this release (selected backend: {backend_name}).",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def _gather_origin_samples(project_dir) -> list:
|
||||
|
|
@ -380,6 +419,9 @@ def cmd_init(args):
|
|||
# Pass 2: detect rooms from folder structure
|
||||
detect_rooms_local(project_dir=args.dir, yes=getattr(args, "yes", False))
|
||||
cfg.init()
|
||||
backend = _backend_arg(args)
|
||||
if backend:
|
||||
cfg.set_backend(backend)
|
||||
|
||||
# Pass 3: protect git repos from accidentally committing per-project files
|
||||
_ensure_mempalace_files_gitignored(args.dir)
|
||||
|
|
@ -615,6 +657,8 @@ def cmd_sync(args):
|
|||
"""Prune drawers whose source files are gitignored, deleted, or moved (#1252)."""
|
||||
from .mcp_server import _wal_log
|
||||
from .palace import MineAlreadyRunning
|
||||
from .backends import detect_backend_for_path
|
||||
from .palace import _backend_artifact_label, resolve_backend_name
|
||||
from .sync import sync_palace
|
||||
|
||||
palace_path = os.path.expanduser(args.palace) if args.palace else MempalaceConfig().palace_path
|
||||
|
|
@ -622,8 +666,16 @@ def cmd_sync(args):
|
|||
if not os.path.isdir(palace_path):
|
||||
print(f"\n No palace found at {palace_path}")
|
||||
return
|
||||
if not os.path.isfile(os.path.join(palace_path, "chroma.sqlite3")):
|
||||
print(f"\n Palace dir at {palace_path} exists but has no chroma.sqlite3 yet.")
|
||||
try:
|
||||
backend_name = resolve_backend_name(palace_path)
|
||||
except Exception as exc: # noqa: BLE001 - user-facing CLI guard
|
||||
print(f"\n Could not resolve palace backend: {exc}", file=sys.stderr)
|
||||
return
|
||||
if detect_backend_for_path(palace_path) is None:
|
||||
print(
|
||||
f"\n Palace dir at {palace_path} exists but has no "
|
||||
f"{_backend_artifact_label(backend_name)} yet."
|
||||
)
|
||||
print(" Run: mempalace mine <dir>")
|
||||
return
|
||||
|
||||
|
|
@ -748,9 +800,11 @@ def cmd_split(args):
|
|||
|
||||
def cmd_migrate(args):
|
||||
"""Migrate palace from a different ChromaDB version."""
|
||||
palace_path = os.path.expanduser(args.palace) if args.palace else MempalaceConfig().palace_path
|
||||
if not _maintenance_requires_chroma(palace_path, "migrate"):
|
||||
raise SystemExit(2)
|
||||
from .migrate import migrate
|
||||
|
||||
palace_path = os.path.expanduser(args.palace) if args.palace else MempalaceConfig().palace_path
|
||||
migrate(
|
||||
palace_path=palace_path,
|
||||
dry_run=args.dry_run,
|
||||
|
|
@ -767,14 +821,24 @@ def cmd_status(args):
|
|||
|
||||
def cmd_repair_status(args):
|
||||
"""Read-only HNSW capacity health check (#1222)."""
|
||||
palace_path = os.path.expanduser(args.palace) if args.palace else MempalaceConfig().palace_path
|
||||
if not _maintenance_requires_chroma(palace_path, "repair-status"):
|
||||
raise SystemExit(2)
|
||||
from .repair import status as repair_status
|
||||
|
||||
palace_path = os.path.expanduser(args.palace) if args.palace else MempalaceConfig().palace_path
|
||||
repair_status(palace_path=palace_path)
|
||||
|
||||
|
||||
def cmd_repair(args):
|
||||
"""Rebuild palace vector index from SQLite metadata."""
|
||||
config = MempalaceConfig()
|
||||
collection_name = config.collection_name
|
||||
palace_path = os.path.abspath(
|
||||
os.path.expanduser(args.palace) if args.palace else config.palace_path
|
||||
)
|
||||
if not _maintenance_requires_chroma(palace_path, "repair"):
|
||||
raise SystemExit(2)
|
||||
|
||||
import shutil
|
||||
from .backends.chroma import ChromaBackend
|
||||
from .migrate import confirm_destructive_action, contains_palace_database
|
||||
|
|
@ -790,12 +854,6 @@ def cmd_repair(args):
|
|||
sqlite_integrity_errors,
|
||||
)
|
||||
|
||||
config = MempalaceConfig()
|
||||
collection_name = config.collection_name
|
||||
palace_path = os.path.abspath(
|
||||
os.path.expanduser(args.palace) if args.palace else config.palace_path
|
||||
)
|
||||
|
||||
if getattr(args, "mode", "legacy") == "max-seq-id":
|
||||
from .repair import repair_max_seq_id
|
||||
|
||||
|
|
@ -1004,12 +1062,15 @@ def cmd_instructions(args):
|
|||
def cmd_mcp(args):
|
||||
"""Show how to wire MemPalace into MCP-capable hosts."""
|
||||
base_server_cmd = "mempalace-mcp"
|
||||
cmd_parts = [base_server_cmd]
|
||||
|
||||
if args.palace:
|
||||
resolved_palace = str(Path(args.palace).expanduser())
|
||||
server_cmd = f"{base_server_cmd} --palace {shlex.quote(resolved_palace)}"
|
||||
else:
|
||||
server_cmd = base_server_cmd
|
||||
cmd_parts.extend(["--palace", shlex.quote(resolved_palace)])
|
||||
backend = _backend_arg(args)
|
||||
if backend:
|
||||
cmd_parts.extend(["--backend", shlex.quote(str(backend).strip().lower())])
|
||||
server_cmd = " ".join(cmd_parts)
|
||||
|
||||
print("MemPalace MCP quick setup:")
|
||||
print(f" claude mcp add mempalace -- {server_cmd}")
|
||||
|
|
@ -1204,12 +1265,23 @@ def main():
|
|||
default=None,
|
||||
help="Where the palace lives (default: from ~/.mempalace/config.json or ~/.mempalace/palace)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
dest="global_backend",
|
||||
default=None,
|
||||
help="Storage backend to use for this command (default: config/env/detected/chroma)",
|
||||
)
|
||||
|
||||
sub = parser.add_subparsers(dest="command")
|
||||
|
||||
# init
|
||||
p_init = sub.add_parser("init", help="Detect rooms from your folder structure")
|
||||
p_init.add_argument("dir", help="Project directory to set up")
|
||||
p_init.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Storage backend to persist for this palace (default: chroma)",
|
||||
)
|
||||
p_init.add_argument(
|
||||
"--yes",
|
||||
action="store_true",
|
||||
|
|
@ -1292,6 +1364,11 @@ def main():
|
|||
# mine
|
||||
p_mine = sub.add_parser("mine", help="Mine files into the palace")
|
||||
p_mine.add_argument("dir", help="Directory to mine")
|
||||
p_mine.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Storage backend to use for this mine (default: config/env/detected/chroma)",
|
||||
)
|
||||
p_mine.add_argument(
|
||||
"--mode",
|
||||
choices=["projects", "convos", "extract"],
|
||||
|
|
@ -1401,6 +1478,11 @@ def main():
|
|||
# search
|
||||
p_search = sub.add_parser("search", help="Find anything, exact words")
|
||||
p_search.add_argument("query", help="What to search for")
|
||||
p_search.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Storage backend to use for this search (default: config/env/detected/chroma)",
|
||||
)
|
||||
p_search.add_argument("--wing", default=None, help="Limit to one project")
|
||||
p_search.add_argument("--room", default=None, help="Limit to one room")
|
||||
p_search.add_argument("--results", type=int, default=5, help="Number of results")
|
||||
|
|
@ -1555,10 +1637,15 @@ def main():
|
|||
)
|
||||
|
||||
# mcp
|
||||
sub.add_parser(
|
||||
p_mcp = sub.add_parser(
|
||||
"mcp",
|
||||
help="Show MCP setup command for connecting MemPalace to your AI client",
|
||||
)
|
||||
p_mcp.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Storage backend to include in the MCP startup command",
|
||||
)
|
||||
|
||||
# status
|
||||
# migrate
|
||||
|
|
@ -1575,9 +1662,15 @@ def main():
|
|||
"--yes", action="store_true", help="Skip confirmation for destructive changes"
|
||||
)
|
||||
|
||||
sub.add_parser("status", help="Show what's been filed")
|
||||
p_status = sub.add_parser("status", help="Show what's been filed")
|
||||
p_status.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Storage backend to use for status (default: config/env/detected/chroma)",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
_apply_backend_arg(args)
|
||||
|
||||
if not args.command:
|
||||
parser.print_help()
|
||||
|
|
|
|||
|
|
@ -191,6 +191,7 @@ def sanitize_content(value: str, max_length: int = 100_000) -> str:
|
|||
|
||||
DEFAULT_PALACE_PATH = os.path.expanduser("~/.mempalace/palace")
|
||||
DEFAULT_COLLECTION_NAME = "mempalace_drawers"
|
||||
DEFAULT_BACKEND = "chroma"
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
|
|
@ -325,6 +326,63 @@ class MempalaceConfig:
|
|||
"""ChromaDB collection name."""
|
||||
return self._file_config.get("collection_name", DEFAULT_COLLECTION_NAME)
|
||||
|
||||
@property
|
||||
def backend(self):
|
||||
"""Storage backend name.
|
||||
|
||||
Read from ``config.json`` first, then ``MEMPALACE_BACKEND``, then
|
||||
``"chroma"`` for backwards compatibility with existing palaces.
|
||||
"""
|
||||
cfg_val = self._file_config.get("backend")
|
||||
if cfg_val:
|
||||
return str(cfg_val).strip().lower()
|
||||
env_val = os.environ.get("MEMPALACE_BACKEND")
|
||||
if env_val:
|
||||
return env_val.strip().lower()
|
||||
return DEFAULT_BACKEND
|
||||
|
||||
@property
|
||||
def qdrant_url(self):
|
||||
"""Qdrant endpoint for the opt-in ``qdrant`` backend.
|
||||
|
||||
Defaults to localhost so selecting Qdrant never silently sends memory
|
||||
to a remote service. Users can point at a LAN or cloud endpoint via
|
||||
config or ``MEMPALACE_QDRANT_URL`` when they deliberately choose that.
|
||||
"""
|
||||
env_val = os.environ.get("MEMPALACE_QDRANT_URL")
|
||||
if env_val:
|
||||
return env_val.strip()
|
||||
return str(self._file_config.get("qdrant_url", "http://localhost:6333")).strip()
|
||||
|
||||
@property
|
||||
def qdrant_api_key(self):
|
||||
"""API key for the opt-in ``qdrant`` backend, if configured."""
|
||||
env_val = os.environ.get("MEMPALACE_QDRANT_API_KEY")
|
||||
if env_val:
|
||||
return env_val
|
||||
value = self._file_config.get("qdrant_api_key")
|
||||
return str(value) if value else None
|
||||
|
||||
@property
|
||||
def qdrant_namespace(self):
|
||||
"""Optional Qdrant collection namespace/prefix."""
|
||||
env_val = os.environ.get("MEMPALACE_QDRANT_NAMESPACE")
|
||||
if env_val:
|
||||
return env_val.strip()
|
||||
value = self._file_config.get("qdrant_namespace")
|
||||
return str(value).strip() if value else None
|
||||
|
||||
@property
|
||||
def qdrant_timeout(self):
|
||||
"""Qdrant HTTP timeout in seconds."""
|
||||
env_val = os.environ.get("MEMPALACE_QDRANT_TIMEOUT")
|
||||
raw = env_val if env_val is not None else self._file_config.get("qdrant_timeout", 10.0)
|
||||
try:
|
||||
timeout = float(raw)
|
||||
except (TypeError, ValueError):
|
||||
timeout = 10.0
|
||||
return timeout if timeout > 0 else 10.0
|
||||
|
||||
@property
|
||||
def people_map(self):
|
||||
"""Mapping of name variants to canonical names."""
|
||||
|
|
@ -560,6 +618,24 @@ class MempalaceConfig:
|
|||
except (OSError, NotImplementedError):
|
||||
pass
|
||||
|
||||
def set_backend(self, backend: str) -> None:
|
||||
"""Persist the storage backend choice to ``config.json``."""
|
||||
backend = str(backend).strip().lower()
|
||||
from .backends import get_backend_class
|
||||
|
||||
get_backend_class(backend)
|
||||
self._file_config["backend"] = backend
|
||||
self._config_dir.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
with open(self._config_file, "w", encoding="utf-8") as f:
|
||||
json.dump(self._file_config, f, indent=2, ensure_ascii=False)
|
||||
except OSError:
|
||||
pass
|
||||
try:
|
||||
self._config_file.chmod(0o600)
|
||||
except (OSError, NotImplementedError):
|
||||
pass
|
||||
|
||||
@property
|
||||
def topic_tunnel_min_count(self):
|
||||
"""Minimum number of overlapping confirmed topics required to create
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ accumulate. This module finds drawers from the same source_file that
|
|||
are too similar (cosine distance < threshold), keeps the longest/richest
|
||||
version, and deletes the rest.
|
||||
|
||||
No API calls — uses ChromaDB's built-in embedding similarity.
|
||||
No API calls — uses the configured local vector backend's similarity search.
|
||||
|
||||
Usage (standalone):
|
||||
python -m mempalace.dedup # dedup all
|
||||
|
|
@ -27,7 +27,7 @@ import os
|
|||
import time
|
||||
from collections import defaultdict
|
||||
|
||||
from .backends.chroma import ChromaBackend
|
||||
from .palace import get_collection
|
||||
|
||||
|
||||
COLLECTION_NAME = "mempalace_drawers"
|
||||
|
|
@ -130,7 +130,7 @@ def dedup_source_group(col, drawer_ids, threshold=DEFAULT_THRESHOLD, dry_run=Tru
|
|||
def show_stats(palace_path=None):
|
||||
"""Show duplication statistics without making changes."""
|
||||
palace_path = palace_path or _get_palace_path()
|
||||
col = ChromaBackend().get_collection(palace_path, COLLECTION_NAME)
|
||||
col = get_collection(palace_path, COLLECTION_NAME)
|
||||
|
||||
groups = get_source_groups(col)
|
||||
|
||||
|
|
@ -162,7 +162,7 @@ def dedup_palace(
|
|||
print(" MemPalace Deduplicator")
|
||||
print(f"{'=' * 55}")
|
||||
|
||||
col = ChromaBackend().get_collection(palace_path, COLLECTION_NAME)
|
||||
col = get_collection(palace_path, COLLECTION_NAME)
|
||||
|
||||
print(f" Palace: {palace_path}")
|
||||
print(f" Drawers: {col.count():,}")
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ from .backends.chroma import ( # noqa: E402
|
|||
_pin_hnsw_threads,
|
||||
hnsw_capacity_status,
|
||||
)
|
||||
from .backends import BackendMismatchError, PalaceRef, detect_backend_for_path # noqa: E402
|
||||
from .query_sanitizer import sanitize_query # noqa: E402
|
||||
from .searcher import search_memories # noqa: E402
|
||||
from .palace_graph import ( # noqa: E402
|
||||
|
|
@ -160,6 +161,11 @@ def _parse_args():
|
|||
metavar="PATH",
|
||||
help="Path to the palace directory (overrides config file and env var)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
metavar="NAME",
|
||||
help="Storage backend to use (default: config/env/detected/chroma)",
|
||||
)
|
||||
args, unknown = parser.parse_known_args()
|
||||
if unknown:
|
||||
logger.debug("Ignoring unknown args: %s", unknown)
|
||||
|
|
@ -170,6 +176,13 @@ _args = _parse_args()
|
|||
|
||||
if _args.palace:
|
||||
os.environ["MEMPALACE_PALACE_PATH"] = os.path.abspath(_args.palace)
|
||||
if _args.backend:
|
||||
backend_name = str(_args.backend).strip().lower()
|
||||
from .backends import get_backend_class # noqa: E402
|
||||
|
||||
get_backend_class(backend_name)
|
||||
os.environ["MEMPALACE_BACKEND_EXPLICIT"] = backend_name
|
||||
os.environ["MEMPALACE_BACKEND"] = backend_name
|
||||
|
||||
_config = MempalaceConfig()
|
||||
|
||||
|
|
@ -289,6 +302,9 @@ def _call_kg(op):
|
|||
|
||||
_client_cache = None
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_collection_open_error = None
|
||||
_palace_db_inode = 0 # inode of chroma.sqlite3 at cache time
|
||||
_palace_db_mtime = 0.0 # mtime of chroma.sqlite3 at cache time
|
||||
|
||||
|
|
@ -313,21 +329,27 @@ def _force_chroma_cache_reset() -> None:
|
|||
global \
|
||||
_client_cache, \
|
||||
_collection_cache, \
|
||||
_collection_cache_backend, \
|
||||
_collection_cache_palace, \
|
||||
_collection_open_error, \
|
||||
_palace_db_inode, \
|
||||
_palace_db_mtime, \
|
||||
_metadata_cache, \
|
||||
_metadata_cache_time
|
||||
_client_cache = None
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_collection_open_error = None
|
||||
_palace_db_inode = 0
|
||||
_palace_db_mtime = 0.0
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
try:
|
||||
from .palace import _DEFAULT_BACKEND
|
||||
from .palace import get_backend_for_palace
|
||||
|
||||
_DEFAULT_BACKEND._clients.pop(_config.palace_path, None)
|
||||
_DEFAULT_BACKEND._freshness.pop(_config.palace_path, None)
|
||||
backend = get_backend_for_palace(_config.palace_path)
|
||||
backend.close_palace(PalaceRef(id=_config.palace_path, local_path=_config.palace_path))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
@ -356,6 +378,11 @@ def _refresh_vector_disabled_flag() -> None:
|
|||
would defeat the point.
|
||||
"""
|
||||
global _vector_disabled, _vector_disabled_reason, _vector_capacity_status
|
||||
if not _is_chroma_backend():
|
||||
_vector_disabled = False
|
||||
_vector_disabled_reason = ""
|
||||
_vector_capacity_status = None
|
||||
return
|
||||
try:
|
||||
info = hnsw_capacity_status(_config.palace_path, _config.collection_name)
|
||||
except Exception:
|
||||
|
|
@ -447,10 +474,15 @@ def _get_client():
|
|||
global \
|
||||
_client_cache, \
|
||||
_collection_cache, \
|
||||
_collection_cache_backend, \
|
||||
_collection_cache_palace, \
|
||||
_collection_open_error, \
|
||||
_palace_db_inode, \
|
||||
_palace_db_mtime, \
|
||||
_metadata_cache, \
|
||||
_metadata_cache_time
|
||||
if not _is_chroma_backend():
|
||||
raise RuntimeError("_get_client is only available for the Chroma backend")
|
||||
db_path = os.path.join(_config.palace_path, "chroma.sqlite3")
|
||||
try:
|
||||
st = os.stat(db_path)
|
||||
|
|
@ -467,6 +499,9 @@ def _get_client():
|
|||
if not os.path.isfile(db_path) and _collection_cache is not None:
|
||||
_client_cache = None
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_collection_open_error = None
|
||||
_palace_db_inode = 0
|
||||
_palace_db_mtime = 0.0
|
||||
# Fall through to normal reconnect which will handle missing DB
|
||||
|
|
@ -482,6 +517,9 @@ def _get_client():
|
|||
_refresh_vector_disabled_flag()
|
||||
_client_cache = ChromaBackend.make_client(_config.palace_path)
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_collection_open_error = None
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
_palace_db_inode = current_inode
|
||||
|
|
@ -490,7 +528,7 @@ def _get_client():
|
|||
|
||||
|
||||
def _get_collection(create=False):
|
||||
"""Return the ChromaDB collection, caching the client between calls.
|
||||
"""Return the configured backend collection, caching handles between calls.
|
||||
|
||||
On failure, log the exception and retry once after clearing the client
|
||||
and collection caches. Tools were silently returning ``None`` when a
|
||||
|
|
@ -501,9 +539,103 @@ def _get_collection(create=False):
|
|||
``quarantine_stale_hnsw`` per #1322), so the second attempt heals the
|
||||
common stale-handle / stale-HNSW case automatically.
|
||||
"""
|
||||
global _client_cache, _collection_cache, _metadata_cache, _metadata_cache_time
|
||||
global \
|
||||
_client_cache, \
|
||||
_collection_cache, \
|
||||
_collection_cache_backend, \
|
||||
_collection_cache_palace, \
|
||||
_collection_open_error, \
|
||||
_palace_db_inode, \
|
||||
_palace_db_mtime, \
|
||||
_metadata_cache, \
|
||||
_metadata_cache_time
|
||||
try:
|
||||
backend_name = _selected_backend_name()
|
||||
except (BackendMismatchError, KeyError) as exc:
|
||||
logger.warning("backend resolution failed for %s: %s", _config.palace_path, exc)
|
||||
_collection_open_error = {
|
||||
"error": "Backend mismatch"
|
||||
if isinstance(exc, BackendMismatchError)
|
||||
else "Unknown backend",
|
||||
"details": str(exc),
|
||||
"hint": "Select the matching backend or use a fresh palace directory.",
|
||||
}
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
return None
|
||||
|
||||
if backend_name != "chroma":
|
||||
for attempt in range(2):
|
||||
try:
|
||||
if (
|
||||
_collection_cache is not None
|
||||
and _collection_cache_backend == backend_name
|
||||
and _collection_cache_palace == _config.palace_path
|
||||
):
|
||||
_collection_open_error = None
|
||||
return _collection_cache
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
if _collection_cache is None:
|
||||
from .palace import get_collection as palace_get_collection
|
||||
|
||||
_collection_cache = palace_get_collection(
|
||||
_config.palace_path,
|
||||
collection_name=_config.collection_name,
|
||||
create=create,
|
||||
backend=backend_name,
|
||||
)
|
||||
_collection_cache_backend = backend_name
|
||||
_collection_cache_palace = _config.palace_path
|
||||
_collection_open_error = None
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
return _collection_cache
|
||||
except (BackendMismatchError, KeyError) as exc:
|
||||
logger.warning("backend open failed for %s: %s", _config.palace_path, exc)
|
||||
_collection_open_error = {
|
||||
"error": "Backend mismatch"
|
||||
if isinstance(exc, BackendMismatchError)
|
||||
else "Unknown backend",
|
||||
"details": str(exc),
|
||||
"hint": "Select the matching backend or use a fresh palace directory.",
|
||||
}
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
return None
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"_get_collection generic attempt %d/2 failed (palace=%s, create=%s)",
|
||||
attempt + 1,
|
||||
_config.palace_path,
|
||||
create,
|
||||
)
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
_collection_open_error = {
|
||||
"error": "Backend open failed",
|
||||
"details": "Could not open the selected backend collection.",
|
||||
"hint": "Run: mempalace status or mempalace repair-status for diagnostics.",
|
||||
}
|
||||
return None
|
||||
|
||||
for attempt in range(2):
|
||||
try:
|
||||
if _collection_cache is not None and (
|
||||
_collection_cache_backend not in (None, "chroma")
|
||||
or _collection_cache_palace not in (None, _config.palace_path)
|
||||
):
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
client = _get_client()
|
||||
# ChromaDB 1.x persists the EF *identity* (its ``name()``) with the
|
||||
# collection but not the EF *instance/configuration*. So a reader or
|
||||
|
|
@ -550,6 +682,9 @@ def _get_collection(create=False):
|
|||
)
|
||||
_pin_hnsw_threads(raw)
|
||||
_collection_cache = ChromaCollection(raw, palace_path=_config.palace_path)
|
||||
_collection_cache_backend = "chroma"
|
||||
_collection_cache_palace = _config.palace_path
|
||||
_collection_open_error = None
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
elif _collection_cache is None:
|
||||
|
|
@ -558,9 +693,29 @@ def _get_collection(create=False):
|
|||
raw = client.get_collection(_config.collection_name, **ef_kwargs)
|
||||
_pin_hnsw_threads(raw)
|
||||
_collection_cache = ChromaCollection(raw, palace_path=_config.palace_path)
|
||||
_collection_cache_backend = "chroma"
|
||||
_collection_cache_palace = _config.palace_path
|
||||
_collection_open_error = None
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
return _collection_cache
|
||||
except (BackendMismatchError, KeyError) as exc:
|
||||
_collection_open_error = {
|
||||
"error": "Backend mismatch"
|
||||
if isinstance(exc, BackendMismatchError)
|
||||
else "Unknown backend",
|
||||
"details": str(exc),
|
||||
"hint": "Select the matching backend or use a fresh palace directory.",
|
||||
}
|
||||
_client_cache = None
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_palace_db_inode = 0
|
||||
_palace_db_mtime = 0.0
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
return None
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"_get_collection attempt %d/2 failed (palace=%s, create=%s)",
|
||||
|
|
@ -575,8 +730,30 @@ def _get_collection(create=False):
|
|||
# collection cleanly, healing the common stale-handle case.
|
||||
_client_cache = None
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_palace_db_inode = 0
|
||||
_palace_db_mtime = 0.0
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
_collection_open_error = {
|
||||
"error": "Backend open failed",
|
||||
"details": "Could not open the Chroma collection.",
|
||||
"hint": "Run: mempalace repair-status for diagnostics.",
|
||||
}
|
||||
_client_cache = None
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_palace_db_inode = 0
|
||||
_palace_db_mtime = 0.0
|
||||
_metadata_cache = None
|
||||
_metadata_cache_time = 0
|
||||
_collection_open_error = _collection_open_error or {
|
||||
"error": "Backend open failed",
|
||||
"details": "Could not open the selected backend collection.",
|
||||
"hint": "Run: mempalace status or mempalace repair-status for diagnostics.",
|
||||
}
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -587,6 +764,42 @@ def _no_palace():
|
|||
}
|
||||
|
||||
|
||||
def _collection_error_or_no_palace():
|
||||
if not _collection_open_error:
|
||||
return _no_palace()
|
||||
result = dict(_collection_open_error)
|
||||
try:
|
||||
result["backend"] = _selected_backend_name()
|
||||
except Exception:
|
||||
pass
|
||||
return result
|
||||
|
||||
|
||||
def _selected_backend_name() -> str:
|
||||
from .palace import resolve_backend_name
|
||||
|
||||
return resolve_backend_name(
|
||||
_config.palace_path,
|
||||
explicit=os.environ.get("MEMPALACE_BACKEND_EXPLICIT"),
|
||||
)
|
||||
|
||||
|
||||
def _is_chroma_backend() -> bool:
|
||||
try:
|
||||
return _selected_backend_name() == "chroma"
|
||||
except Exception:
|
||||
logger.debug("backend resolution failed", exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
def _backend_db_exists() -> bool:
|
||||
try:
|
||||
return detect_backend_for_path(_config.palace_path) is not None
|
||||
except Exception:
|
||||
logger.debug("backend artifact detection failed", exc_info=True)
|
||||
return False
|
||||
|
||||
|
||||
# ==================== HELPERS ====================
|
||||
|
||||
|
||||
|
|
@ -722,6 +935,7 @@ def _tool_status_via_sqlite() -> dict:
|
|||
"rooms": rooms,
|
||||
"protocol": PALACE_PROTOCOL,
|
||||
"aaak_dialect": AAAK_SPEC,
|
||||
"backend": "chroma",
|
||||
"vector_disabled": True,
|
||||
"vector_disabled_reason": _vector_disabled_reason,
|
||||
}
|
||||
|
|
@ -739,7 +953,7 @@ def tool_status():
|
|||
# #1222 failure mode, opening the persistent client to call .count()
|
||||
# can segfault — short-circuit to a pure-sqlite path when divergence
|
||||
# is detected so status stays reachable.
|
||||
db_exists = os.path.isfile(os.path.join(_config.palace_path, "chroma.sqlite3"))
|
||||
db_exists = _backend_db_exists()
|
||||
_refresh_vector_disabled_flag()
|
||||
|
||||
if _vector_disabled:
|
||||
|
|
@ -750,7 +964,7 @@ def tool_status():
|
|||
# accidentally creating a palace in a non-existent directory (#830).
|
||||
col = _get_collection(create=db_exists)
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
count = col.count()
|
||||
wings = {}
|
||||
rooms = {}
|
||||
|
|
@ -760,6 +974,7 @@ def tool_status():
|
|||
"rooms": rooms,
|
||||
"protocol": PALACE_PROTOCOL,
|
||||
"aaak_dialect": AAAK_SPEC,
|
||||
"backend": _selected_backend_name(),
|
||||
}
|
||||
try:
|
||||
all_meta = _get_cached_metadata(col)
|
||||
|
|
@ -812,7 +1027,7 @@ When WRITING AAAK: use entity codes, mark emotions, keep structure tight."""
|
|||
def tool_list_wings():
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
wings = {}
|
||||
result = {"wings": wings}
|
||||
try:
|
||||
|
|
@ -835,7 +1050,7 @@ def tool_list_rooms(wing: str = None):
|
|||
return {"error": str(e)}
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
rooms = {}
|
||||
result = {"wing": wing or "all", "rooms": rooms}
|
||||
try:
|
||||
|
|
@ -855,7 +1070,7 @@ def tool_list_rooms(wing: str = None):
|
|||
def tool_get_taxonomy():
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
taxonomy = {}
|
||||
result = {"taxonomy": taxonomy}
|
||||
try:
|
||||
|
|
@ -925,6 +1140,7 @@ def tool_search(
|
|||
n_results=limit,
|
||||
max_distance=dist,
|
||||
vector_disabled=_vector_disabled,
|
||||
collection_name=_config.collection_name,
|
||||
)
|
||||
if not _is_transient_index_error(result):
|
||||
result["index_recovered"] = True
|
||||
|
|
@ -963,7 +1179,7 @@ def tool_check_duplicate(content: str, threshold: float = 0.9):
|
|||
}
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
try:
|
||||
content = strip_lone_surrogates(content)
|
||||
results = col.query(
|
||||
|
|
@ -1009,7 +1225,7 @@ def tool_traverse_graph(start_room: str, max_hops: int = 2):
|
|||
max_hops = max(1, min(max_hops, 10))
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
return traverse(start_room, col=col, max_hops=max_hops)
|
||||
|
||||
|
||||
|
|
@ -1022,7 +1238,7 @@ def tool_find_tunnels(wing_a: str = None, wing_b: str = None):
|
|||
return {"error": str(e)}
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
return find_tunnels(wing_a, wing_b, col=col)
|
||||
|
||||
|
||||
|
|
@ -1030,7 +1246,7 @@ def tool_graph_stats():
|
|||
"""Palace graph overview: nodes, tunnels, edges, connectivity."""
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
return graph_stats(col=col)
|
||||
|
||||
|
||||
|
|
@ -1096,6 +1312,8 @@ def tool_follow_tunnels(wing: str, room: str):
|
|||
except ValueError as e:
|
||||
return {"error": str(e)}
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _collection_error_or_no_palace()
|
||||
return follow_tunnels(wing, room, col=col)
|
||||
|
||||
|
||||
|
|
@ -1130,7 +1348,7 @@ def tool_add_drawer(
|
|||
|
||||
col = _get_collection(create=True)
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
|
||||
drawer_id = (
|
||||
f"drawer_{wing}_{room}_{hashlib.sha256((wing + room + content).encode()).hexdigest()[:24]}"
|
||||
|
|
@ -1244,7 +1462,7 @@ def tool_delete_drawer(drawer_id: str):
|
|||
global _metadata_cache
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
existing = col.get(ids=[drawer_id])
|
||||
if not existing["ids"]:
|
||||
return {"success": False, "error": f"Drawer not found: {drawer_id}"}
|
||||
|
|
@ -1314,7 +1532,7 @@ def tool_get_drawer(drawer_id: str):
|
|||
"""Fetch a single drawer by ID. Returns full content and metadata."""
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
try:
|
||||
result = col.get(ids=[drawer_id], include=["documents", "metadatas"])
|
||||
if not result["ids"]:
|
||||
|
|
@ -1352,7 +1570,7 @@ def tool_list_drawers(wing: str = None, room: str = None, limit: int = 20, offse
|
|||
return {"error": str(e)}
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
try:
|
||||
where = None
|
||||
conditions = []
|
||||
|
|
@ -1409,7 +1627,7 @@ def tool_update_drawer(drawer_id: str, content: str = None, wing: str = None, ro
|
|||
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
try:
|
||||
existing = col.get(ids=[drawer_id], include=["documents", "metadatas"])
|
||||
if not existing["ids"]:
|
||||
|
|
@ -1627,7 +1845,7 @@ def tool_diary_write(agent_name: str, entry: str, topic: str = "general", wing:
|
|||
room = "diary"
|
||||
col = _get_collection(create=True)
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
|
||||
now = datetime.now()
|
||||
entry_id = (
|
||||
|
|
@ -1750,7 +1968,7 @@ def tool_diary_read(agent_name: str, last_n: int = 10, wing: str = ""):
|
|||
last_n = max(1, min(last_n, 100))
|
||||
col = _get_collection()
|
||||
if not col:
|
||||
return _no_palace()
|
||||
return _collection_error_or_no_palace()
|
||||
|
||||
# Build filter: always scope by agent + room=diary. Wing is optional —
|
||||
# when empty, return entries across all wings for this agent (matches
|
||||
|
|
@ -1884,6 +2102,9 @@ def tool_reconnect():
|
|||
global \
|
||||
_client_cache, \
|
||||
_collection_cache, \
|
||||
_collection_cache_backend, \
|
||||
_collection_cache_palace, \
|
||||
_collection_open_error, \
|
||||
_palace_db_inode, \
|
||||
_palace_db_mtime, \
|
||||
_vector_disabled, \
|
||||
|
|
@ -1891,29 +2112,60 @@ def tool_reconnect():
|
|||
from . import palace as palace_module
|
||||
|
||||
close_errors = []
|
||||
palace_ref = PalaceRef(id=_config.palace_path, local_path=_config.palace_path)
|
||||
closed_backend_names = set()
|
||||
cached_backend_name = _collection_cache_backend
|
||||
try:
|
||||
palace_module._DEFAULT_BACKEND.close_palace(_config.palace_path)
|
||||
backend = palace_module.get_backend_for_palace(_config.palace_path)
|
||||
backend.close_palace(palace_ref)
|
||||
if getattr(backend, "name", None):
|
||||
closed_backend_names.add(backend.name)
|
||||
except Exception as exc:
|
||||
logger.debug("Failed to close shared palace backend during reconnect", exc_info=True)
|
||||
close_errors.append(f"backend close_palace failed: {exc}")
|
||||
try:
|
||||
from chromadb.api.client import SharedSystemClient
|
||||
if cached_backend_name and cached_backend_name not in closed_backend_names:
|
||||
try:
|
||||
from .backends import get_backend
|
||||
|
||||
clear_system_cache = getattr(SharedSystemClient, "clear_system_cache", None)
|
||||
if callable(clear_system_cache):
|
||||
clear_system_cache()
|
||||
else:
|
||||
get_backend(cached_backend_name).close_palace(palace_ref)
|
||||
closed_backend_names.add(cached_backend_name)
|
||||
except Exception as exc:
|
||||
logger.debug(
|
||||
"SharedSystemClient.clear_system_cache is unavailable; skipping shared Chroma cache clear during reconnect"
|
||||
"Failed to close previously cached %s backend during reconnect",
|
||||
cached_backend_name,
|
||||
exc_info=True,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug(
|
||||
"Failed to clear Chroma shared system cache during reconnect",
|
||||
exc_info=True,
|
||||
)
|
||||
close_errors.append(f"shared Chroma cache clear failed: {exc}")
|
||||
close_errors.append(f"cached {cached_backend_name} close_palace failed: {exc}")
|
||||
if _client_cache is not None:
|
||||
try:
|
||||
close = getattr(_client_cache, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
except Exception as exc:
|
||||
logger.debug("Failed to close MCP-local Chroma client during reconnect", exc_info=True)
|
||||
close_errors.append(f"local Chroma client close failed: {exc}")
|
||||
if _is_chroma_backend():
|
||||
try:
|
||||
from chromadb.api.client import SharedSystemClient
|
||||
|
||||
clear_system_cache = getattr(SharedSystemClient, "clear_system_cache", None)
|
||||
if callable(clear_system_cache):
|
||||
clear_system_cache()
|
||||
else:
|
||||
logger.debug(
|
||||
"SharedSystemClient.clear_system_cache is unavailable; skipping shared Chroma cache clear during reconnect"
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug(
|
||||
"Failed to clear Chroma shared system cache during reconnect",
|
||||
exc_info=True,
|
||||
)
|
||||
close_errors.append(f"shared Chroma cache clear failed: {exc}")
|
||||
_client_cache = None
|
||||
_collection_cache = None
|
||||
_collection_cache_backend = None
|
||||
_collection_cache_palace = None
|
||||
_collection_open_error = None
|
||||
_palace_db_inode = 0
|
||||
_palace_db_mtime = 0.0
|
||||
# Force probe re-run on next _get_client by clearing the flag now;
|
||||
|
|
@ -1933,12 +2185,17 @@ def tool_reconnect():
|
|||
try:
|
||||
col = _get_collection()
|
||||
if col is None:
|
||||
open_error = _collection_error_or_no_palace()
|
||||
result = {
|
||||
"success": False,
|
||||
"message": "No palace found after reconnect",
|
||||
"message": open_error.get("error", "No palace found after reconnect"),
|
||||
"drawers": 0,
|
||||
"vector_disabled": _vector_disabled,
|
||||
}
|
||||
if "details" in open_error:
|
||||
result["details"] = open_error["details"]
|
||||
if "hint" in open_error:
|
||||
result["hint"] = open_error["hint"]
|
||||
if close_errors:
|
||||
result["error"] = "; ".join(close_errors)
|
||||
return result
|
||||
|
|
@ -2724,8 +2981,17 @@ def _maybe_eager_warmup_embedder() -> None:
|
|||
)
|
||||
return
|
||||
palace_path = _config.palace_path
|
||||
db_path = os.path.join(palace_path, "chroma.sqlite3")
|
||||
if not os.path.isfile(db_path):
|
||||
try:
|
||||
backend_name = _selected_backend_name()
|
||||
except Exception as exc: # fail-soft per docstring
|
||||
logger.warning(
|
||||
"MEMPALACE_EAGER_WARMUP=%s: backend resolution failed for %s (%s)",
|
||||
raw,
|
||||
palace_path,
|
||||
exc,
|
||||
)
|
||||
return
|
||||
if not _backend_db_exists():
|
||||
# Pre-check (NOT a try/except on _ChromaNotFoundError, which never
|
||||
# propagates out of _get_collection — see docstring). No palace
|
||||
# file means nothing to warm AND avoids the chromadb-client
|
||||
|
|
@ -2769,9 +3035,11 @@ def _maybe_eager_warmup_embedder() -> None:
|
|||
type(exc).__name__,
|
||||
)
|
||||
else:
|
||||
warmed = "embedder + HNSW ready" if backend_name == "chroma" else "embedder + backend ready"
|
||||
logger.info(
|
||||
"MEMPALACE_EAGER_WARMUP=%s: embedder + HNSW ready (palace=%s, device=%s)",
|
||||
"MEMPALACE_EAGER_WARMUP=%s: %s (palace=%s, device=%s)",
|
||||
raw,
|
||||
warmed,
|
||||
palace_path,
|
||||
device,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -13,8 +13,19 @@ import sys
|
|||
import threading
|
||||
from typing import Optional
|
||||
|
||||
from .backends import BackendClosedError, CollectionNotInitializedError, PalaceNotFoundError
|
||||
from .backends.chroma import ChromaBackend
|
||||
from .backends import (
|
||||
BackendClosedError,
|
||||
BackendMismatchError,
|
||||
CollectionNotInitializedError,
|
||||
PalaceNotFoundError,
|
||||
PalaceRef,
|
||||
detect_backend_for_path,
|
||||
detect_backends_for_path,
|
||||
get_backend,
|
||||
get_backend_class,
|
||||
resolve_backend_for_palace,
|
||||
)
|
||||
from .backends.embedding_wrapper import EmbeddingCollection
|
||||
from .entity_detector import _apply_known_systems_prepass, _get_coca_filter
|
||||
|
||||
logger = logging.getLogger("mempalace_mcp")
|
||||
|
|
@ -45,7 +56,8 @@ SKIP_DIRS = {
|
|||
"target",
|
||||
}
|
||||
|
||||
_DEFAULT_BACKEND = ChromaBackend()
|
||||
_DEFAULT_BACKEND = get_backend("chroma")
|
||||
_EXPLICIT_BACKEND_ENV = "MEMPALACE_BACKEND_EXPLICIT"
|
||||
|
||||
# Schema version for drawer normalization. Bump when the normalization
|
||||
# pipeline changes in a way that existing drawers should be rebuilt to pick up
|
||||
|
|
@ -62,22 +74,120 @@ def get_collection(
|
|||
palace_path: str,
|
||||
collection_name: Optional[str] = None,
|
||||
create: bool = True,
|
||||
backend: Optional[str] = None,
|
||||
):
|
||||
"""Get the palace collection through the backend layer."""
|
||||
if collection_name is None:
|
||||
from .config import get_configured_collection_name
|
||||
|
||||
collection_name = get_configured_collection_name()
|
||||
return _DEFAULT_BACKEND.get_collection(
|
||||
backend_obj = get_backend_for_palace(palace_path, explicit=backend)
|
||||
palace_ref = PalaceRef(id=palace_path, local_path=palace_path)
|
||||
try:
|
||||
collection = backend_obj.get_collection(
|
||||
palace=palace_ref,
|
||||
collection_name=collection_name,
|
||||
create=create,
|
||||
)
|
||||
except TypeError as exc:
|
||||
if "unexpected keyword argument 'palace'" not in str(exc):
|
||||
raise
|
||||
collection = backend_obj.get_collection(
|
||||
palace_path,
|
||||
collection_name=collection_name,
|
||||
create=create,
|
||||
)
|
||||
if "requires_explicit_embeddings" in getattr(backend_obj, "capabilities", frozenset()):
|
||||
return EmbeddingCollection(collection)
|
||||
return collection
|
||||
|
||||
|
||||
def get_closets_collection(
|
||||
palace_path: str,
|
||||
create: bool = True,
|
||||
backend: Optional[str] = None,
|
||||
):
|
||||
"""Get the closets collection — the searchable index layer."""
|
||||
return get_collection(
|
||||
palace_path,
|
||||
collection_name=collection_name,
|
||||
collection_name="mempalace_closets",
|
||||
create=create,
|
||||
backend=backend,
|
||||
)
|
||||
|
||||
|
||||
def get_closets_collection(palace_path: str, create: bool = True):
|
||||
"""Get the closets collection — the searchable index layer."""
|
||||
return get_collection(palace_path, collection_name="mempalace_closets", create=create)
|
||||
def _config_backend_value(palace_path: str) -> Optional[str]:
|
||||
try:
|
||||
from .config import MempalaceConfig
|
||||
|
||||
cfg = MempalaceConfig()
|
||||
cfg_palace = os.path.abspath(os.path.expanduser(cfg.palace_path))
|
||||
target_palace = os.path.abspath(os.path.expanduser(palace_path))
|
||||
if cfg_palace != target_palace:
|
||||
return None
|
||||
value = cfg._file_config.get("backend")
|
||||
return str(value).strip().lower() if value else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _env_backend_value() -> Optional[str]:
|
||||
value = os.environ.get("MEMPALACE_BACKEND")
|
||||
return value.strip().lower() if value else None
|
||||
|
||||
|
||||
def resolve_backend_name(palace_path: str, explicit: Optional[str] = None) -> str:
|
||||
"""Resolve and validate the selected backend for ``palace_path``.
|
||||
|
||||
Public resolution order:
|
||||
|
||||
1. Explicit CLI/MCP flag or direct ``get_collection(..., backend=...)``.
|
||||
2. ``backend`` in ``~/.mempalace/config.json``.
|
||||
3. ``MEMPALACE_BACKEND``.
|
||||
4. Detected existing palace artifacts.
|
||||
5. ``chroma``.
|
||||
|
||||
If artifacts for a different backend are already present, raise
|
||||
``BackendMismatchError`` so normal write paths cannot silently mix storage
|
||||
formats in one palace directory.
|
||||
"""
|
||||
explicit = explicit or os.environ.get(_EXPLICIT_BACKEND_ENV)
|
||||
selected = resolve_backend_for_palace(
|
||||
explicit=explicit.strip().lower() if explicit else None,
|
||||
config_value=_config_backend_value(palace_path),
|
||||
env_value=_env_backend_value(),
|
||||
palace_path=palace_path,
|
||||
default="chroma",
|
||||
)
|
||||
get_backend_class(selected)
|
||||
detected_backends = detect_backends_for_path(palace_path)
|
||||
if len(detected_backends) > 1:
|
||||
raise BackendMismatchError(
|
||||
f"palace at {palace_path!r} contains multiple backend artifacts: "
|
||||
f"{', '.join(detected_backends)}"
|
||||
)
|
||||
detected = detected_backends[0] if detected_backends else None
|
||||
if detected and detected != selected:
|
||||
raise BackendMismatchError(
|
||||
f"palace at {palace_path!r} contains {detected!r} backend artifacts, "
|
||||
f"but {selected!r} was selected"
|
||||
)
|
||||
return selected
|
||||
|
||||
|
||||
def get_backend_for_palace(palace_path: str, explicit: Optional[str] = None):
|
||||
"""Return the resolved backend instance for ``palace_path``."""
|
||||
return get_backend(resolve_backend_name(palace_path, explicit=explicit))
|
||||
|
||||
|
||||
def _backend_artifact_label(backend_name: Optional[str]) -> str:
|
||||
if backend_name == "chroma":
|
||||
return "chroma.sqlite3"
|
||||
if backend_name == "qdrant":
|
||||
return "qdrant_backend.json"
|
||||
if backend_name == "sqlite_exact":
|
||||
return "sqlite_exact.sqlite3"
|
||||
return "backend database"
|
||||
|
||||
|
||||
def _open_collection_or_explain(
|
||||
|
|
@ -85,6 +195,7 @@ def _open_collection_or_explain(
|
|||
*,
|
||||
collection_name: Optional[str] = None,
|
||||
out=None,
|
||||
opener=None,
|
||||
):
|
||||
"""Open the palace collection or print a state-specific message and return ``None``.
|
||||
|
||||
|
|
@ -101,11 +212,11 @@ def _open_collection_or_explain(
|
|||
first when the vector path is disabled (see PR #831 / issue #830).
|
||||
|
||||
State A: palace dir is absent.
|
||||
State B: dir is present but ``chroma.sqlite3`` is absent. The helper
|
||||
short-circuits to a message before reaching the backend, because
|
||||
``chromadb.PersistentClient`` lazily creates the DB file on first
|
||||
open — calling the backend on this state would silently mutate
|
||||
the filesystem for what should be a read-only inspection.
|
||||
State B: dir is present but no backend database artifact is present.
|
||||
The helper short-circuits to a message before reaching the backend,
|
||||
because some backends lazily create their DB file on first open —
|
||||
calling the backend on this state would silently mutate the filesystem
|
||||
for what should be a read-only inspection.
|
||||
State C: DB is present but the ``mempalace_drawers`` collection has
|
||||
never been bootstrapped (``init`` ran, ``mine`` has not).
|
||||
State D: healthy — returns the opened collection.
|
||||
|
|
@ -116,17 +227,33 @@ def _open_collection_or_explain(
|
|||
callable (e.g. a repair progress emitter) to route messages through it.
|
||||
"""
|
||||
emit = out if out is not None else print
|
||||
open_collection = opener or get_collection
|
||||
|
||||
if not os.path.isdir(palace_path):
|
||||
emit(f"\n No palace found at {palace_path}")
|
||||
emit(" Run: mempalace init <dir> then mempalace mine <dir>")
|
||||
return None
|
||||
if not os.path.isfile(os.path.join(palace_path, "chroma.sqlite3")):
|
||||
emit(f"\n Palace dir at {palace_path} exists but has no chroma.sqlite3 yet.")
|
||||
try:
|
||||
backend_name = resolve_backend_name(palace_path)
|
||||
except BackendMismatchError as e:
|
||||
emit(f"\n Backend mismatch at {palace_path}: {e}")
|
||||
emit(" Select the matching backend or use a fresh palace directory.")
|
||||
return None
|
||||
detected = detect_backend_for_path(palace_path)
|
||||
if detected is None:
|
||||
emit(
|
||||
f"\n Palace dir at {palace_path} exists but has no "
|
||||
f"{_backend_artifact_label(backend_name)} yet."
|
||||
)
|
||||
emit(" Run: mempalace mine <dir>")
|
||||
return None
|
||||
try:
|
||||
return get_collection(palace_path, collection_name=collection_name, create=False)
|
||||
return open_collection(
|
||||
palace_path,
|
||||
collection_name=collection_name,
|
||||
create=False,
|
||||
backend=backend_name,
|
||||
)
|
||||
except CollectionNotInitializedError:
|
||||
emit(f"\n Palace at {palace_path} is initialized but empty (no drawers yet).")
|
||||
emit(" Run: mempalace mine <dir>")
|
||||
|
|
@ -135,6 +262,10 @@ def _open_collection_or_explain(
|
|||
emit(f"\n No palace found at {palace_path}")
|
||||
emit(" Run: mempalace init <dir> then mempalace mine <dir>")
|
||||
return None
|
||||
except BackendMismatchError as e:
|
||||
emit(f"\n Backend mismatch at {palace_path}: {e}")
|
||||
emit(" Select the matching backend or use a fresh palace directory.")
|
||||
return None
|
||||
except BackendClosedError:
|
||||
# Surface this as a programmer error, not a palace-state UX message:
|
||||
# a closed backend means the caller violated the backend lifecycle,
|
||||
|
|
@ -484,6 +615,9 @@ def _validate_palace_fts5_after_mine(palace_path: str) -> None:
|
|||
operator sees the same recovery banner regardless of which command surfaces
|
||||
the bug.
|
||||
"""
|
||||
if resolve_backend_name(palace_path) != "chroma":
|
||||
return
|
||||
|
||||
# Defer-import: keeps the repair module graph out of mine's hot import path.
|
||||
from .repair import _close_chroma_handles, sqlite_integrity_errors
|
||||
|
||||
|
|
|
|||
|
|
@ -16,8 +16,19 @@ import re
|
|||
import sqlite3
|
||||
from pathlib import Path
|
||||
|
||||
from .backends import CollectionNotInitializedError, PalaceNotFoundError
|
||||
from .palace import get_closets_collection, get_collection
|
||||
from .backends import (
|
||||
BackendError,
|
||||
BackendMismatchError,
|
||||
CollectionNotInitializedError,
|
||||
PalaceNotFoundError,
|
||||
UnsupportedCapabilityError,
|
||||
)
|
||||
from .palace import (
|
||||
_open_collection_or_explain,
|
||||
get_closets_collection,
|
||||
get_collection,
|
||||
resolve_backend_name,
|
||||
)
|
||||
|
||||
# Closet pointer line format: "topic|entities|→drawer_id_a,drawer_id_b"
|
||||
# Multiple lines may join with newlines inside one closet document.
|
||||
|
|
@ -296,32 +307,11 @@ def search(query: str, palace_path: str, wing: str = None, room: str = None, n_r
|
|||
Search the palace. Returns verbatim drawer content.
|
||||
Optionally filter by wing (project) or room (aspect).
|
||||
"""
|
||||
# Filesystem-first checks distinguish State A / State B before reaching
|
||||
# chromadb. PersistentClient lazily creates chroma.sqlite3 on first open
|
||||
# of an empty palace dir, so without these checks State B collapses into
|
||||
# the "initialized but empty" State C message and mutates the dir as a
|
||||
# side effect of a read-only search call (#1498).
|
||||
if not os.path.isdir(palace_path):
|
||||
print(f"\n No palace found at {palace_path}")
|
||||
print(" Run: mempalace init <dir> then mempalace mine <dir>")
|
||||
raise SearchError(f"No palace found at {palace_path}")
|
||||
if not os.path.isfile(os.path.join(palace_path, "chroma.sqlite3")):
|
||||
print(f"\n Palace dir at {palace_path} exists but has no chroma.sqlite3 yet.")
|
||||
print(" Run: mempalace mine <dir>")
|
||||
col = _open_collection_or_explain(palace_path, opener=get_collection)
|
||||
if col is None:
|
||||
if not os.path.isdir(palace_path):
|
||||
raise SearchError(f"No palace found at {palace_path}")
|
||||
raise SearchError(f"No palace database at {palace_path}")
|
||||
try:
|
||||
col = get_collection(palace_path, create=False)
|
||||
except CollectionNotInitializedError as e:
|
||||
# State C from #1498: palace initialized but never mined.
|
||||
print(f"\n Palace at {palace_path} is initialized but empty (no drawers yet).")
|
||||
print(" Run: mempalace mine <dir>")
|
||||
raise SearchError(f"Palace at {palace_path} is initialized but empty") from e
|
||||
except PalaceNotFoundError as e:
|
||||
# Backend filesystem-race fallback: dir was deleted between our
|
||||
# check above and the backend call. Same message as State A.
|
||||
print(f"\n No palace found at {palace_path}")
|
||||
print(" Run: mempalace init <dir> then mempalace mine <dir>")
|
||||
raise SearchError(f"No palace found at {palace_path}") from e
|
||||
|
||||
# Alert the user if this palace predates hnsw:space=cosine being set on
|
||||
# creation — their similarity scores will be junk until they run repair.
|
||||
|
|
@ -636,14 +626,14 @@ def _bm25_only_via_sqlite(
|
|||
|
||||
def _merge_bm25_union_candidates(
|
||||
hits: list,
|
||||
drawers_col,
|
||||
query: str,
|
||||
palace_path: str,
|
||||
wing: str,
|
||||
room: str,
|
||||
n_results: int,
|
||||
max_distance: float = 0.0,
|
||||
) -> None:
|
||||
"""Append top-K BM25-only candidates from sqlite into ``hits`` in place.
|
||||
"""Append top-K backend lexical candidates into ``hits`` in place.
|
||||
|
||||
Used by ``search_memories(..., candidate_strategy="union")`` to widen
|
||||
the rerank pool's *source* (not just its size) — vector-only candidate
|
||||
|
|
@ -668,19 +658,41 @@ def _merge_bm25_union_candidates(
|
|||
if max_distance > 0.0:
|
||||
return
|
||||
|
||||
where = build_where_filter(wing, room)
|
||||
try:
|
||||
bm25_extra = _bm25_only_via_sqlite(
|
||||
query,
|
||||
palace_path,
|
||||
wing=wing,
|
||||
room=room,
|
||||
lexical = drawers_col.lexical_search(
|
||||
query=query,
|
||||
n_results=n_results * 3,
|
||||
_include_internal=True,
|
||||
).get("results", [])
|
||||
where=where or None,
|
||||
)
|
||||
except UnsupportedCapabilityError:
|
||||
raise
|
||||
except Exception:
|
||||
logger.debug("candidate_strategy=union: BM25 fetch failed", exc_info=True)
|
||||
logger.debug("candidate_strategy=union: lexical fetch failed", exc_info=True)
|
||||
return
|
||||
|
||||
bm25_extra = []
|
||||
for hit in lexical.hits:
|
||||
meta = hit.metadata or {}
|
||||
full_source = meta.get("source_file", "") or ""
|
||||
bm25_extra.append(
|
||||
{
|
||||
"text": hit.document or "",
|
||||
"wing": meta.get("wing", "unknown"),
|
||||
"room": meta.get("room", "unknown"),
|
||||
"source_file": Path(full_source).name if full_source else "?",
|
||||
"created_at": meta.get("filed_at", "unknown"),
|
||||
"similarity": None,
|
||||
"distance": None,
|
||||
"effective_distance": None,
|
||||
"closet_boost": 0.0,
|
||||
"matched_via": "bm25_backend",
|
||||
"bm25_score": round(float(hit.score), 3),
|
||||
"_source_file_full": full_source,
|
||||
"_chunk_index": meta.get("chunk_index"),
|
||||
}
|
||||
)
|
||||
|
||||
def _dedup_key(entry: dict):
|
||||
full = entry.get("_source_file_full")
|
||||
ci = entry.get("_chunk_index")
|
||||
|
|
@ -728,8 +740,8 @@ def _validate_candidate_strategy(strategy: str) -> None:
|
|||
def _apply_candidate_strategy(
|
||||
strategy: str,
|
||||
hits: list,
|
||||
drawers_col,
|
||||
query: str,
|
||||
palace_path: str,
|
||||
wing: str,
|
||||
room: str,
|
||||
n_results: int,
|
||||
|
|
@ -742,7 +754,120 @@ def _apply_candidate_strategy(
|
|||
"""
|
||||
merger = _CANDIDATE_MERGERS[strategy]
|
||||
if merger is not None:
|
||||
merger(hits, query, palace_path, wing, room, n_results, max_distance=max_distance)
|
||||
merger(hits, drawers_col, query, wing, room, n_results, max_distance=max_distance)
|
||||
|
||||
|
||||
def _finalize_candidate_hits(
|
||||
*,
|
||||
candidate_strategy: str,
|
||||
hits: list,
|
||||
drawers_col,
|
||||
query: str,
|
||||
wing: str,
|
||||
room: str,
|
||||
n_results: int,
|
||||
max_distance: float,
|
||||
) -> tuple:
|
||||
try:
|
||||
_apply_candidate_strategy(
|
||||
candidate_strategy,
|
||||
hits,
|
||||
drawers_col,
|
||||
query,
|
||||
wing,
|
||||
room,
|
||||
n_results,
|
||||
max_distance=max_distance,
|
||||
)
|
||||
except UnsupportedCapabilityError:
|
||||
return [], {
|
||||
"error": "candidate_strategy='union' requires a backend with lexical_search support",
|
||||
"unsupported_capability": "supports_lexical_search",
|
||||
"hint": "Use candidate_strategy='vector' or select a backend that supports lexical search.",
|
||||
}
|
||||
|
||||
hits = _hybrid_rank(hits, query)[:n_results]
|
||||
for h in hits:
|
||||
h.pop("_sort_key", None)
|
||||
h.pop("_source_file_full", None)
|
||||
h.pop("_chunk_index", None)
|
||||
return hits, None
|
||||
|
||||
|
||||
def _backend_mismatch_result(error: BackendMismatchError) -> dict:
|
||||
return {
|
||||
"error": "Backend mismatch",
|
||||
"details": str(error),
|
||||
"hint": "Select the matching backend or use a fresh palace directory.",
|
||||
}
|
||||
|
||||
|
||||
def _unknown_backend_result(error: KeyError) -> dict:
|
||||
return {
|
||||
"error": "Unknown backend",
|
||||
"details": str(error),
|
||||
"hint": "Check MEMPALACE_BACKEND or the configured backend name.",
|
||||
}
|
||||
|
||||
|
||||
def _vector_disabled_search(
|
||||
*,
|
||||
query: str,
|
||||
palace_path: str,
|
||||
wing: str,
|
||||
room: str,
|
||||
n_results: int,
|
||||
collection_name: str,
|
||||
) -> dict:
|
||||
try:
|
||||
backend_name = resolve_backend_name(palace_path)
|
||||
except BackendMismatchError as e:
|
||||
return _backend_mismatch_result(e)
|
||||
except KeyError as e:
|
||||
return _unknown_backend_result(e)
|
||||
if backend_name != "chroma":
|
||||
return {
|
||||
"error": "vector_disabled fallback is Chroma-only",
|
||||
"unsupported_capability": "chroma_hnsw_fallback",
|
||||
"backend": backend_name,
|
||||
"hint": "Disable vector_disabled for non-Chroma backends.",
|
||||
}
|
||||
return _bm25_only_via_sqlite(
|
||||
query,
|
||||
palace_path,
|
||||
wing=wing,
|
||||
room=room,
|
||||
n_results=n_results,
|
||||
collection_name=collection_name,
|
||||
)
|
||||
|
||||
|
||||
def _open_search_collection(palace_path: str, collection_name: str):
|
||||
try:
|
||||
return get_collection(palace_path, collection_name=collection_name, create=False), None
|
||||
except BackendMismatchError as e:
|
||||
return None, _backend_mismatch_result(e)
|
||||
except KeyError as e:
|
||||
return None, _unknown_backend_result(e)
|
||||
except (CollectionNotInitializedError, PalaceNotFoundError) as e:
|
||||
logger.error("No palace found at %s: %s", palace_path, e)
|
||||
return None, {
|
||||
"error": "No palace found",
|
||||
"hint": "Run: mempalace init <dir> && mempalace mine <dir>",
|
||||
}
|
||||
except BackendError as e:
|
||||
logger.error("Backend error opening palace at %s: %s", palace_path, e)
|
||||
return None, {
|
||||
"error": "Backend error",
|
||||
"details": str(e),
|
||||
"hint": "Check the selected backend configuration and availability.",
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error("No palace found at %s: %s", palace_path, e)
|
||||
return None, {
|
||||
"error": "No palace found",
|
||||
"hint": "Run: mempalace init <dir> && mempalace mine <dir>",
|
||||
}
|
||||
|
||||
|
||||
def search_memories(
|
||||
|
|
@ -780,14 +905,12 @@ def search_memories(
|
|||
``n_results * 3`` rows from the vector index are the rerank pool.
|
||||
Cheap; works well when query and target docs agree in the
|
||||
embedding space.
|
||||
* ``"union"`` — also pull top ``n_results * 3`` BM25 candidates
|
||||
from the sqlite FTS5 index and merge them into the rerank pool
|
||||
(deduped by source_file). Catches docs with strong BM25 signal
|
||||
that are vector-distant from the query (e.g. terminology guides
|
||||
looked up by narrative-shaped queries; policy clauses surfaced
|
||||
by scenario descriptions). Adds one sqlite open + FTS5 MATCH
|
||||
per query; perf cost is small but unmeasured at corpus scale.
|
||||
Opt in until the cost is characterized.
|
||||
* ``"union"`` — also pull top ``n_results * 3`` lexical candidates
|
||||
through the backend's ``lexical_search`` capability and merge
|
||||
them into the rerank pool (deduped by source_file). Catches docs
|
||||
with strong BM25 signal that are vector-distant from the query.
|
||||
Perf depends on the selected backend; opt in until the cost is
|
||||
characterized.
|
||||
|
||||
When ``max_distance > 0.0`` is also set, BM25-only candidates
|
||||
are skipped — they have no vector distance and would silently
|
||||
|
|
@ -799,23 +922,18 @@ def search_memories(
|
|||
_validate_candidate_strategy(candidate_strategy)
|
||||
|
||||
if vector_disabled:
|
||||
return _bm25_only_via_sqlite(
|
||||
query,
|
||||
palace_path,
|
||||
return _vector_disabled_search(
|
||||
query=query,
|
||||
palace_path=palace_path,
|
||||
wing=wing,
|
||||
room=room,
|
||||
n_results=n_results,
|
||||
collection_name=collection_name,
|
||||
)
|
||||
|
||||
try:
|
||||
drawers_col = get_collection(palace_path, collection_name=collection_name, create=False)
|
||||
except Exception as e:
|
||||
logger.error("No palace found at %s: %s", palace_path, e)
|
||||
return {
|
||||
"error": "No palace found",
|
||||
"hint": "Run: mempalace init <dir> && mempalace mine <dir>",
|
||||
}
|
||||
drawers_col, open_error = _open_search_collection(palace_path, collection_name)
|
||||
if open_error:
|
||||
return open_error
|
||||
|
||||
where = build_where_filter(wing, room)
|
||||
|
||||
|
|
@ -985,31 +1103,24 @@ def search_memories(
|
|||
|
||||
# Candidate strategy hook: optionally widen the rerank pool's *source*
|
||||
# before ranking. Default ("vector") is a no-op; "union" merges top-K
|
||||
# BM25 candidates from sqlite. See `_apply_candidate_strategy`.
|
||||
# backend lexical candidates. See `_apply_candidate_strategy`.
|
||||
# ``max_distance`` is forwarded so union mode can refuse to inject
|
||||
# BM25-only (distance=None) candidates that would silently bypass the
|
||||
# caller's strict distance threshold.
|
||||
_apply_candidate_strategy(
|
||||
candidate_strategy,
|
||||
hits,
|
||||
query,
|
||||
palace_path,
|
||||
wing,
|
||||
room,
|
||||
n_results,
|
||||
# The helper also runs the final BM25 hybrid re-rank and strips internal
|
||||
# dedup fields before returning.
|
||||
hits, strategy_error = _finalize_candidate_hits(
|
||||
candidate_strategy=candidate_strategy,
|
||||
hits=hits,
|
||||
drawers_col=drawers_col,
|
||||
query=query,
|
||||
wing=wing,
|
||||
room=room,
|
||||
n_results=n_results,
|
||||
max_distance=max_distance,
|
||||
)
|
||||
|
||||
# BM25 hybrid re-rank within the final candidate set, then trim back
|
||||
# to the requested size. Without the trim, ``candidate_strategy="union"``
|
||||
# would return up to 4× ``n_results`` (vector hits + BM25 union pool),
|
||||
# breaking the existing ``search_memories`` size contract that the MCP
|
||||
# ``limit`` parameter is built on.
|
||||
hits = _hybrid_rank(hits, query)[:n_results]
|
||||
for h in hits:
|
||||
h.pop("_sort_key", None)
|
||||
h.pop("_source_file_full", None)
|
||||
h.pop("_chunk_index", None)
|
||||
if strategy_error:
|
||||
return strategy_error
|
||||
|
||||
return {
|
||||
"query": query,
|
||||
|
|
|
|||
|
|
@ -60,6 +60,8 @@ mempalace-mcp = "mempalace.mcp_server:main"
|
|||
|
||||
[project.entry-points."mempalace.backends"]
|
||||
chroma = "mempalace.backends.chroma:ChromaBackend"
|
||||
qdrant = "mempalace.backends.qdrant:QdrantBackend"
|
||||
sqlite_exact = "mempalace.backends.sqlite_exact:SQLiteExactBackend"
|
||||
|
||||
# RFC 002 source-adapter entry-point group. Core publishes no first-party
|
||||
# adapters under this group yet; ``miner.py`` and ``convo_miner.py`` migrate
|
||||
|
|
|
|||
|
|
@ -79,6 +79,12 @@ def _reset_mcp_cache():
|
|||
|
||||
mcp_server._client_cache = None
|
||||
mcp_server._collection_cache = None
|
||||
if hasattr(mcp_server, "_collection_cache_backend"):
|
||||
mcp_server._collection_cache_backend = None
|
||||
if hasattr(mcp_server, "_collection_cache_palace"):
|
||||
mcp_server._collection_cache_palace = None
|
||||
if hasattr(mcp_server, "_collection_open_error"):
|
||||
mcp_server._collection_open_error = None
|
||||
except AttributeError:
|
||||
pass
|
||||
|
||||
|
|
|
|||
|
|
@ -181,6 +181,66 @@ def test_chroma_detect_matches_palace_with_chroma_sqlite(tmp_path):
|
|||
assert ChromaBackend.detect(str(tmp_path.parent)) is False
|
||||
|
||||
|
||||
def test_chroma_lexical_search_uses_sqlite_fts_not_full_collection_scan(tmp_path):
|
||||
db_path = tmp_path / "chroma.sqlite3"
|
||||
conn = sqlite3.connect(db_path)
|
||||
conn.executescript(
|
||||
"""
|
||||
CREATE TABLE collections (id INTEGER PRIMARY KEY, name TEXT NOT NULL);
|
||||
CREATE TABLE segments (id INTEGER PRIMARY KEY, collection INTEGER NOT NULL);
|
||||
CREATE TABLE embeddings (id INTEGER PRIMARY KEY, segment_id INTEGER NOT NULL, created_at TEXT);
|
||||
CREATE TABLE embedding_metadata (
|
||||
id INTEGER,
|
||||
key TEXT,
|
||||
string_value TEXT,
|
||||
int_value INTEGER,
|
||||
float_value REAL,
|
||||
bool_value INTEGER
|
||||
);
|
||||
CREATE VIRTUAL TABLE embedding_fulltext_search USING fts5(string_value);
|
||||
"""
|
||||
)
|
||||
conn.execute("INSERT INTO collections(id, name) VALUES (1, 'mempalace_drawers')")
|
||||
conn.execute("INSERT INTO segments(id, collection) VALUES (1, 1)")
|
||||
ids = list(range(1, 14))
|
||||
for emb_id in ids:
|
||||
wing = "target" if emb_id == 13 else "old"
|
||||
doc = "needle shared lexical note"
|
||||
conn.execute(
|
||||
"INSERT INTO embeddings(id, segment_id, created_at) VALUES (?, 1, ?)",
|
||||
(emb_id, f"2026-01-01T00:00:{emb_id:02d}"),
|
||||
)
|
||||
conn.execute(
|
||||
"INSERT INTO embedding_fulltext_search(rowid, string_value) VALUES (?, ?)",
|
||||
(emb_id, doc),
|
||||
)
|
||||
conn.execute(
|
||||
"INSERT INTO embedding_metadata(id, key, string_value) VALUES (?, 'chroma:document', ?)",
|
||||
(emb_id, doc),
|
||||
)
|
||||
conn.execute(
|
||||
"INSERT INTO embedding_metadata(id, key, string_value) VALUES (?, 'wing', ?)",
|
||||
(emb_id, wing),
|
||||
)
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
class _NoScanCollection:
|
||||
name = "mempalace_drawers"
|
||||
|
||||
def count(self):
|
||||
raise AssertionError("lexical_search should use Chroma sqlite FTS")
|
||||
|
||||
def get(self, **_kwargs):
|
||||
raise AssertionError("lexical_search should use Chroma sqlite FTS")
|
||||
|
||||
collection = ChromaCollection(_NoScanCollection(), palace_path=str(tmp_path))
|
||||
|
||||
hits = collection.lexical_search(query="needle", n_results=1, where={"wing": "target"}).hits
|
||||
|
||||
assert [hit.metadata["wing"] for hit in hits] == ["target"]
|
||||
|
||||
|
||||
def test_query_rejects_missing_input():
|
||||
fake = _FakeCollection()
|
||||
collection = ChromaCollection(fake)
|
||||
|
|
|
|||
|
|
@ -713,6 +713,40 @@ def test_main_status_dispatches():
|
|||
mock_cmd.assert_called_once()
|
||||
|
||||
|
||||
def test_main_backend_flag_sets_explicit_backend(monkeypatch):
|
||||
monkeypatch.delenv("MEMPALACE_BACKEND_EXPLICIT", raising=False)
|
||||
monkeypatch.delenv("MEMPALACE_BACKEND", raising=False)
|
||||
with (
|
||||
patch("sys.argv", ["mempalace", "status", "--backend", "sqlite_exact"]),
|
||||
patch("mempalace.cli.cmd_status") as mock_cmd,
|
||||
):
|
||||
main()
|
||||
|
||||
mock_cmd.assert_called_once()
|
||||
args = mock_cmd.call_args.args[0]
|
||||
assert args.backend == "sqlite_exact"
|
||||
assert os.environ["MEMPALACE_BACKEND_EXPLICIT"] == "sqlite_exact"
|
||||
os.environ.pop("MEMPALACE_BACKEND_EXPLICIT", None)
|
||||
os.environ.pop("MEMPALACE_BACKEND", None)
|
||||
|
||||
|
||||
def test_main_backend_flag_accepts_qdrant(monkeypatch):
|
||||
monkeypatch.delenv("MEMPALACE_BACKEND_EXPLICIT", raising=False)
|
||||
monkeypatch.delenv("MEMPALACE_BACKEND", raising=False)
|
||||
with (
|
||||
patch("sys.argv", ["mempalace", "status", "--backend", "qdrant"]),
|
||||
patch("mempalace.cli.cmd_status") as mock_cmd,
|
||||
):
|
||||
main()
|
||||
|
||||
mock_cmd.assert_called_once()
|
||||
args = mock_cmd.call_args.args[0]
|
||||
assert args.backend == "qdrant"
|
||||
assert os.environ["MEMPALACE_BACKEND_EXPLICIT"] == "qdrant"
|
||||
os.environ.pop("MEMPALACE_BACKEND_EXPLICIT", None)
|
||||
os.environ.pop("MEMPALACE_BACKEND", None)
|
||||
|
||||
|
||||
def test_main_search_dispatches():
|
||||
with (
|
||||
patch("sys.argv", ["mempalace", "search", "my query"]),
|
||||
|
|
@ -790,6 +824,34 @@ def test_mcp_command_uses_custom_palace_path_when_provided(monkeypatch, capsys):
|
|||
assert captured.err == ""
|
||||
|
||||
|
||||
def test_mcp_command_includes_backend_when_provided(monkeypatch, capsys):
|
||||
monkeypatch.delenv("MEMPALACE_BACKEND_EXPLICIT", raising=False)
|
||||
monkeypatch.delenv("MEMPALACE_BACKEND", raising=False)
|
||||
monkeypatch.setattr(sys, "argv", ["mempalace", "mcp", "--backend", "sqlite_exact"])
|
||||
|
||||
main()
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert "mempalace-mcp --backend sqlite_exact" in captured.out
|
||||
assert captured.err == ""
|
||||
os.environ.pop("MEMPALACE_BACKEND_EXPLICIT", None)
|
||||
os.environ.pop("MEMPALACE_BACKEND", None)
|
||||
|
||||
|
||||
def test_mcp_command_includes_qdrant_backend(monkeypatch, capsys):
|
||||
monkeypatch.delenv("MEMPALACE_BACKEND_EXPLICIT", raising=False)
|
||||
monkeypatch.delenv("MEMPALACE_BACKEND", raising=False)
|
||||
monkeypatch.setattr(sys, "argv", ["mempalace", "mcp", "--backend", "qdrant"])
|
||||
|
||||
main()
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert "mempalace-mcp --backend qdrant" in captured.out
|
||||
assert captured.err == ""
|
||||
os.environ.pop("MEMPALACE_BACKEND_EXPLICIT", None)
|
||||
os.environ.pop("MEMPALACE_BACKEND", None)
|
||||
|
||||
|
||||
def test_main_hook_no_subcommand_prints_help(capsys):
|
||||
with patch("sys.argv", ["mempalace", "hook"]):
|
||||
main()
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ def test_default_config():
|
|||
cfg = MempalaceConfig(config_dir=tempfile.mkdtemp())
|
||||
assert "palace" in cfg.palace_path
|
||||
assert cfg.collection_name == "mempalace_drawers"
|
||||
assert cfg.backend == "chroma"
|
||||
|
||||
|
||||
def test_config_from_file():
|
||||
|
|
@ -27,6 +28,54 @@ def test_config_from_file():
|
|||
assert cfg.palace_path == "/custom/palace"
|
||||
|
||||
|
||||
def test_backend_from_config_wins_over_env(tmp_path, monkeypatch):
|
||||
with open(tmp_path / "config.json", "w") as f:
|
||||
json.dump({"backend": "sqlite_exact"}, f)
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND", "chroma")
|
||||
|
||||
cfg = MempalaceConfig(config_dir=str(tmp_path))
|
||||
assert cfg.backend == "sqlite_exact"
|
||||
|
||||
|
||||
def test_backend_from_env_when_config_absent(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND", "SQLite_Exact")
|
||||
|
||||
cfg = MempalaceConfig(config_dir=str(tmp_path))
|
||||
assert cfg.backend == "sqlite_exact"
|
||||
|
||||
|
||||
def test_qdrant_config_from_env_and_file(tmp_path, monkeypatch):
|
||||
with open(tmp_path / "config.json", "w") as f:
|
||||
json.dump(
|
||||
{
|
||||
"qdrant_url": "http://config.example:6333",
|
||||
"qdrant_api_key": "config-key",
|
||||
"qdrant_namespace": "config-ns",
|
||||
"qdrant_timeout": 2,
|
||||
},
|
||||
f,
|
||||
)
|
||||
monkeypatch.setenv("MEMPALACE_QDRANT_URL", "http://env.example:6333")
|
||||
monkeypatch.setenv("MEMPALACE_QDRANT_API_KEY", "env-key")
|
||||
monkeypatch.setenv("MEMPALACE_QDRANT_NAMESPACE", "env-ns")
|
||||
monkeypatch.setenv("MEMPALACE_QDRANT_TIMEOUT", "3.5")
|
||||
|
||||
cfg = MempalaceConfig(config_dir=str(tmp_path))
|
||||
|
||||
assert cfg.qdrant_url == "http://env.example:6333"
|
||||
assert cfg.qdrant_api_key == "env-key"
|
||||
assert cfg.qdrant_namespace == "env-ns"
|
||||
assert cfg.qdrant_timeout == 3.5
|
||||
|
||||
|
||||
def test_set_backend_persists_choice(tmp_path):
|
||||
cfg = MempalaceConfig(config_dir=str(tmp_path))
|
||||
cfg.set_backend("sqlite_exact")
|
||||
|
||||
reloaded = MempalaceConfig(config_dir=str(tmp_path))
|
||||
assert reloaded.backend == "sqlite_exact"
|
||||
|
||||
|
||||
def test_embedding_device_defaults_to_auto(monkeypatch):
|
||||
monkeypatch.delenv("MEMPALACE_EMBEDDING_DEVICE", raising=False)
|
||||
cfg = MempalaceConfig(config_dir=tempfile.mkdtemp())
|
||||
|
|
@ -115,6 +164,17 @@ def test_init():
|
|||
cfg = MempalaceConfig(config_dir=tmpdir)
|
||||
cfg.init()
|
||||
assert os.path.exists(os.path.join(tmpdir, "config.json"))
|
||||
with open(os.path.join(tmpdir, "config.json")) as f:
|
||||
saved = json.load(f)
|
||||
assert "backend" not in saved
|
||||
assert MempalaceConfig(config_dir=tmpdir).backend == "chroma"
|
||||
|
||||
|
||||
def test_set_backend_rejects_unknown_backend(tmp_path):
|
||||
cfg = MempalaceConfig(config_dir=str(tmp_path))
|
||||
|
||||
with pytest.raises(KeyError):
|
||||
cfg.set_backend("does_not_exist")
|
||||
|
||||
|
||||
# --- normalize_wing_name ---
|
||||
|
|
|
|||
|
|
@ -198,15 +198,13 @@ def test_dedup_source_group_query_failure_keeps():
|
|||
# ── show_stats ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _install_mock_backend(mock_backend_cls, collection):
|
||||
mock_backend = MagicMock()
|
||||
mock_backend.get_collection.return_value = collection
|
||||
mock_backend_cls.return_value = mock_backend
|
||||
return mock_backend
|
||||
def _install_mock_collection(mock_get_collection, collection):
|
||||
mock_get_collection.return_value = collection
|
||||
return collection
|
||||
|
||||
|
||||
@patch("mempalace.dedup.ChromaBackend")
|
||||
def test_show_stats(mock_backend_cls, tmp_path):
|
||||
@patch("mempalace.dedup.get_collection")
|
||||
def test_show_stats(mock_get_collection, tmp_path):
|
||||
mock_col = MagicMock()
|
||||
mock_col.count.return_value = 5
|
||||
mock_col.get.side_effect = [
|
||||
|
|
@ -222,7 +220,7 @@ def test_show_stats(mock_backend_cls, tmp_path):
|
|||
},
|
||||
{"ids": []},
|
||||
]
|
||||
_install_mock_backend(mock_backend_cls, mock_col)
|
||||
_install_mock_collection(mock_get_collection, mock_col)
|
||||
|
||||
dedup.show_stats(palace_path=str(tmp_path)) # should not raise
|
||||
|
||||
|
|
@ -232,11 +230,11 @@ def test_show_stats(mock_backend_cls, tmp_path):
|
|||
|
||||
@patch("mempalace.dedup.dedup_source_group")
|
||||
@patch("mempalace.dedup.get_source_groups")
|
||||
@patch("mempalace.dedup.ChromaBackend")
|
||||
def test_dedup_palace_dry_run(mock_backend_cls, mock_groups, mock_dedup_group, tmp_path):
|
||||
@patch("mempalace.dedup.get_collection")
|
||||
def test_dedup_palace_dry_run(mock_get_collection, mock_groups, mock_dedup_group, tmp_path):
|
||||
mock_col = MagicMock()
|
||||
mock_col.count.return_value = 10
|
||||
_install_mock_backend(mock_backend_cls, mock_col)
|
||||
_install_mock_collection(mock_get_collection, mock_col)
|
||||
|
||||
mock_groups.return_value = {"a.txt": ["d1", "d2", "d3", "d4", "d5"]}
|
||||
mock_dedup_group.return_value = (["d1", "d2", "d3"], ["d4", "d5"])
|
||||
|
|
@ -247,11 +245,11 @@ def test_dedup_palace_dry_run(mock_backend_cls, mock_groups, mock_dedup_group, t
|
|||
|
||||
@patch("mempalace.dedup.dedup_source_group")
|
||||
@patch("mempalace.dedup.get_source_groups")
|
||||
@patch("mempalace.dedup.ChromaBackend")
|
||||
def test_dedup_palace_with_wing(mock_backend_cls, mock_groups, mock_dedup_group, tmp_path):
|
||||
@patch("mempalace.dedup.get_collection")
|
||||
def test_dedup_palace_with_wing(mock_get_collection, mock_groups, mock_dedup_group, tmp_path):
|
||||
mock_col = MagicMock()
|
||||
mock_col.count.return_value = 10
|
||||
_install_mock_backend(mock_backend_cls, mock_col)
|
||||
_install_mock_collection(mock_get_collection, mock_col)
|
||||
|
||||
mock_groups.return_value = {}
|
||||
dedup.dedup_palace(palace_path=str(tmp_path), wing="test_wing", dry_run=True)
|
||||
|
|
@ -260,11 +258,11 @@ def test_dedup_palace_with_wing(mock_backend_cls, mock_groups, mock_dedup_group,
|
|||
|
||||
@patch("mempalace.dedup.dedup_source_group")
|
||||
@patch("mempalace.dedup.get_source_groups")
|
||||
@patch("mempalace.dedup.ChromaBackend")
|
||||
def test_dedup_palace_no_groups(mock_backend_cls, mock_groups, mock_dedup_group, tmp_path):
|
||||
@patch("mempalace.dedup.get_collection")
|
||||
def test_dedup_palace_no_groups(mock_get_collection, mock_groups, mock_dedup_group, tmp_path):
|
||||
mock_col = MagicMock()
|
||||
mock_col.count.return_value = 3
|
||||
_install_mock_backend(mock_backend_cls, mock_col)
|
||||
_install_mock_collection(mock_get_collection, mock_col)
|
||||
|
||||
mock_groups.return_value = {}
|
||||
dedup.dedup_palace(palace_path=str(tmp_path), dry_run=True)
|
||||
|
|
|
|||
|
|
@ -695,6 +695,75 @@ class TestReadTools:
|
|||
assert "project" in result["wings"]
|
||||
assert "notes" in result["wings"]
|
||||
|
||||
def test_status_sqlite_exact_backend_has_no_hnsw_fields(
|
||||
self, monkeypatch, config, palace_path, kg
|
||||
):
|
||||
import mempalace.backends.embedding_wrapper as embedding_wrapper
|
||||
from mempalace.palace import get_collection
|
||||
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "sqlite_exact")
|
||||
monkeypatch.setattr(
|
||||
embedding_wrapper,
|
||||
"_embed_texts",
|
||||
lambda texts: [[float(len(text)), 1.0] for text in texts],
|
||||
)
|
||||
col = get_collection(palace_path, create=True)
|
||||
col.add(
|
||||
ids=["drawer_sqlite"],
|
||||
documents=["verbatim sqlite drawer"],
|
||||
metadatas=[{"wing": "w", "room": "r"}],
|
||||
)
|
||||
|
||||
_patch_mcp_server(monkeypatch, config, kg)
|
||||
from mempalace import mcp_server
|
||||
|
||||
monkeypatch.setattr(mcp_server, "_collection_cache", None)
|
||||
result = mcp_server.tool_status()
|
||||
|
||||
assert result["backend"] == "sqlite_exact"
|
||||
assert result["total_drawers"] == 1
|
||||
assert "hnsw_capacity" not in result
|
||||
assert result.get("vector_disabled") is not True
|
||||
|
||||
def test_status_qdrant_backend_has_no_hnsw_fields(self, monkeypatch, config, palace_path, kg):
|
||||
from mempalace.backends import GetResult
|
||||
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "qdrant")
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND", "qdrant")
|
||||
with open(os.path.join(palace_path, "qdrant_backend.json"), "w", encoding="utf-8") as f:
|
||||
json.dump({"backend": "qdrant"}, f)
|
||||
|
||||
_patch_mcp_server(monkeypatch, config, kg)
|
||||
from mempalace import mcp_server
|
||||
|
||||
class _FakeQdrantCollection:
|
||||
def count(self):
|
||||
return 2
|
||||
|
||||
def get(self, **_kwargs):
|
||||
return GetResult(
|
||||
ids=["q1", "q2"],
|
||||
documents=[],
|
||||
metadatas=[
|
||||
{"wing": "project", "room": "backend"},
|
||||
{"wing": "project", "room": "api"},
|
||||
],
|
||||
)
|
||||
|
||||
monkeypatch.setattr(mcp_server, "_collection_cache", None)
|
||||
monkeypatch.setattr(mcp_server, "_metadata_cache", None)
|
||||
monkeypatch.setattr(
|
||||
mcp_server, "_get_collection", lambda create=False: _FakeQdrantCollection()
|
||||
)
|
||||
|
||||
result = mcp_server.tool_status()
|
||||
|
||||
assert result["backend"] == "qdrant"
|
||||
assert result["total_drawers"] == 2
|
||||
assert result["wings"] == {"project": 2}
|
||||
assert "hnsw_capacity" not in result
|
||||
assert result.get("vector_disabled") is not True
|
||||
|
||||
def test_status_handles_none_metadata_without_partial(
|
||||
self, monkeypatch, config, palace_path, kg
|
||||
):
|
||||
|
|
@ -2290,7 +2359,69 @@ class TestCacheInvalidation:
|
|||
|
||||
result = mcp_server.tool_reconnect()
|
||||
assert result["success"] is True
|
||||
close_palace.assert_called_once_with(config.palace_path)
|
||||
closed_ref = close_palace.call_args.args[0]
|
||||
assert closed_ref.local_path == config.palace_path
|
||||
|
||||
def test_reconnect_closes_selected_non_chroma_backend(
|
||||
self, monkeypatch, config, palace_path, kg
|
||||
):
|
||||
_patch_mcp_server(monkeypatch, config, kg)
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "sqlite_exact")
|
||||
from mempalace import mcp_server, palace
|
||||
|
||||
closed = []
|
||||
|
||||
class _FakeBackend:
|
||||
def close_palace(self, path):
|
||||
closed.append(path)
|
||||
|
||||
class _FakeCol:
|
||||
def count(self):
|
||||
return 3
|
||||
|
||||
monkeypatch.setattr(palace, "get_backend_for_palace", lambda _path: _FakeBackend())
|
||||
monkeypatch.setattr(mcp_server, "_is_chroma_backend", lambda: False)
|
||||
monkeypatch.setattr(mcp_server, "_get_collection", lambda create=False: _FakeCol())
|
||||
|
||||
result = mcp_server.tool_reconnect()
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["drawers"] == 3
|
||||
assert len(closed) == 1
|
||||
assert closed[0].local_path == palace_path
|
||||
|
||||
def test_reconnect_closes_previously_cached_backend(self, monkeypatch, config, palace_path, kg):
|
||||
_patch_mcp_server(monkeypatch, config, kg)
|
||||
from mempalace import backends, mcp_server, palace
|
||||
|
||||
closed = []
|
||||
|
||||
class _SelectedBackend:
|
||||
name = "sqlite_exact"
|
||||
|
||||
def close_palace(self, ref):
|
||||
closed.append(("selected", ref.local_path))
|
||||
|
||||
class _CachedBackend:
|
||||
name = "chroma"
|
||||
|
||||
def close_palace(self, ref):
|
||||
closed.append(("cached", ref.local_path))
|
||||
|
||||
class _FakeCol:
|
||||
def count(self):
|
||||
return 3
|
||||
|
||||
monkeypatch.setattr(palace, "get_backend_for_palace", lambda _path: _SelectedBackend())
|
||||
monkeypatch.setattr(backends, "get_backend", lambda _name: _CachedBackend())
|
||||
monkeypatch.setattr(mcp_server, "_collection_cache_backend", "chroma")
|
||||
monkeypatch.setattr(mcp_server, "_is_chroma_backend", lambda: False)
|
||||
monkeypatch.setattr(mcp_server, "_get_collection", lambda create=False: _FakeCol())
|
||||
|
||||
result = mcp_server.tool_reconnect()
|
||||
|
||||
assert result["success"] is True
|
||||
assert closed == [("selected", palace_path), ("cached", palace_path)]
|
||||
|
||||
def test_get_collection_create_true_avoids_get_or_create_on_reopen(
|
||||
self, monkeypatch, config, palace_path, kg
|
||||
|
|
|
|||
|
|
@ -0,0 +1,462 @@
|
|||
import os
|
||||
import uuid
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from mempalace.backends import (
|
||||
BackendError,
|
||||
BackendMismatchError,
|
||||
CollectionNotInitializedError,
|
||||
DimensionMismatchError,
|
||||
PalaceRef,
|
||||
available_backends,
|
||||
)
|
||||
from mempalace.backends.qdrant import QdrantBackend
|
||||
|
||||
|
||||
def _get_payload_value(payload, key):
|
||||
value = payload
|
||||
for part in key.split("."):
|
||||
if not isinstance(value, dict):
|
||||
return None
|
||||
value = value.get(part)
|
||||
return value
|
||||
|
||||
|
||||
def _fake_match_condition(point, condition):
|
||||
if "must" in condition or "must_not" in condition or "should" in condition:
|
||||
return _fake_match_filter(point, condition)
|
||||
if "has_id" in condition:
|
||||
return point["id"] in set(condition["has_id"])
|
||||
key = condition.get("key")
|
||||
actual = _get_payload_value(point.get("payload") or {}, key)
|
||||
if "match" in condition:
|
||||
match = condition["match"]
|
||||
if "value" in match:
|
||||
return actual == match["value"]
|
||||
if "any" in match:
|
||||
return actual in set(match["any"] or [])
|
||||
if "text_any" in match:
|
||||
haystack = str(actual or "").lower()
|
||||
return any(token in haystack for token in str(match["text_any"]).lower().split())
|
||||
if "range" in condition:
|
||||
range_spec = condition["range"]
|
||||
try:
|
||||
if "gt" in range_spec and not actual > range_spec["gt"]:
|
||||
return False
|
||||
if "gte" in range_spec and not actual >= range_spec["gte"]:
|
||||
return False
|
||||
if "lt" in range_spec and not actual < range_spec["lt"]:
|
||||
return False
|
||||
if "lte" in range_spec and not actual <= range_spec["lte"]:
|
||||
return False
|
||||
except TypeError:
|
||||
return False
|
||||
return True
|
||||
return True
|
||||
|
||||
|
||||
def _fake_match_filter(point, qdrant_filter):
|
||||
if not qdrant_filter:
|
||||
return True
|
||||
must = qdrant_filter.get("must") or []
|
||||
must_not = qdrant_filter.get("must_not") or []
|
||||
should = qdrant_filter.get("should") or []
|
||||
if any(not _fake_match_condition(point, condition) for condition in must):
|
||||
return False
|
||||
if any(_fake_match_condition(point, condition) for condition in must_not):
|
||||
return False
|
||||
if should and not any(_fake_match_condition(point, condition) for condition in should):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
class _FakeQdrantClient:
|
||||
instances = []
|
||||
|
||||
def __init__(self, _config):
|
||||
self.collections = {}
|
||||
self.query_calls = []
|
||||
self.created_indexes = []
|
||||
_FakeQdrantClient.instances.append(self)
|
||||
|
||||
def request(self, *_args, **_kwargs):
|
||||
return {"result": {}}
|
||||
|
||||
def collection_exists(self, collection):
|
||||
return collection in self.collections
|
||||
|
||||
def get_collection_info(self, collection):
|
||||
if collection not in self.collections:
|
||||
raise AssertionError("collection missing")
|
||||
return {
|
||||
"result": {
|
||||
"config": {
|
||||
"params": {
|
||||
"vectors": {
|
||||
"size": self.collections[collection]["dimension"],
|
||||
"distance": "Cosine",
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
def create_collection(self, collection, dimension):
|
||||
self.collections.setdefault(collection, {"dimension": dimension, "points": {}})
|
||||
|
||||
def create_payload_index(self, collection, field_name, field_schema):
|
||||
self.created_indexes.append((collection, field_name, field_schema))
|
||||
|
||||
def upsert_points(self, collection, points):
|
||||
self.collections.setdefault(
|
||||
collection,
|
||||
{"dimension": len(points[0]["vector"]) if points else 0, "points": {}},
|
||||
)
|
||||
for point in points:
|
||||
self.collections[collection]["points"][point["id"]] = dict(point)
|
||||
|
||||
def query_points(self, collection, *, vector, limit, qdrant_filter, with_vector):
|
||||
self.query_calls.append(qdrant_filter)
|
||||
points = list(self.collections.get(collection, {"points": {}})["points"].values())
|
||||
points = [point for point in points if _fake_match_filter(point, qdrant_filter)]
|
||||
q = np.asarray(vector, dtype=np.float32)
|
||||
scored = []
|
||||
for point in points:
|
||||
vec = np.asarray(point["vector"], dtype=np.float32)
|
||||
denom = float(np.linalg.norm(q)) * float(np.linalg.norm(vec))
|
||||
score = 0.0 if denom <= 0 else float(np.dot(q, vec) / denom)
|
||||
out = {"id": point["id"], "payload": point["payload"], "score": score}
|
||||
if with_vector:
|
||||
out["vector"] = point["vector"]
|
||||
scored.append(out)
|
||||
scored.sort(key=lambda point: point["score"], reverse=True)
|
||||
return scored[:limit]
|
||||
|
||||
def scroll_points(
|
||||
self,
|
||||
collection,
|
||||
*,
|
||||
qdrant_filter=None,
|
||||
limit=256,
|
||||
offset=None,
|
||||
with_vector=False,
|
||||
):
|
||||
points = list(self.collections.get(collection, {"points": {}})["points"].values())
|
||||
points = [point for point in points if _fake_match_filter(point, qdrant_filter)]
|
||||
start = int(offset or 0)
|
||||
selected = points[start : start + limit]
|
||||
next_offset = start + limit if start + limit < len(points) else None
|
||||
out = []
|
||||
for point in selected:
|
||||
item = {"id": point["id"], "payload": point["payload"]}
|
||||
if with_vector:
|
||||
item["vector"] = point["vector"]
|
||||
out.append(item)
|
||||
return out, next_offset
|
||||
|
||||
def delete_points(self, collection, *, point_ids=None, qdrant_filter=None):
|
||||
points = self.collections.get(collection, {"points": {}})["points"]
|
||||
if point_ids is not None:
|
||||
for point_id in point_ids:
|
||||
points.pop(point_id, None)
|
||||
return
|
||||
for point_id, point in list(points.items()):
|
||||
if _fake_match_filter(point, qdrant_filter):
|
||||
points.pop(point_id, None)
|
||||
|
||||
def count_points(self, collection):
|
||||
return len(self.collections.get(collection, {"points": {}})["points"])
|
||||
|
||||
def delete_collection(self, collection):
|
||||
self.collections.pop(collection, None)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_qdrant(monkeypatch):
|
||||
import mempalace.backends.qdrant as qdrant
|
||||
|
||||
_FakeQdrantClient.instances.clear()
|
||||
monkeypatch.setattr(qdrant, "_QdrantRESTClient", _FakeQdrantClient)
|
||||
monkeypatch.delenv("MEMPALACE_QDRANT_URL", raising=False)
|
||||
monkeypatch.delenv("MEMPALACE_QDRANT_API_KEY", raising=False)
|
||||
monkeypatch.delenv("MEMPALACE_QDRANT_NAMESPACE", raising=False)
|
||||
monkeypatch.delenv("MEMPALACE_QDRANT_TIMEOUT", raising=False)
|
||||
return _FakeQdrantClient
|
||||
|
||||
|
||||
def _collection(tmp_path, name="drawers"):
|
||||
backend = QdrantBackend()
|
||||
palace = PalaceRef(id=str(tmp_path), local_path=str(tmp_path))
|
||||
return backend, backend.get_collection(palace=palace, collection_name=name, create=True)
|
||||
|
||||
|
||||
def test_registry_exposes_qdrant():
|
||||
assert "qdrant" in available_backends()
|
||||
|
||||
|
||||
def test_qdrant_add_query_filters_lexical_and_marker(tmp_path, fake_qdrant):
|
||||
backend, col = _collection(tmp_path)
|
||||
assert not os.path.isfile(tmp_path / "qdrant_backend.json")
|
||||
|
||||
col.add(
|
||||
ids=["a", "b", "c"],
|
||||
documents=[
|
||||
"alpha backend note",
|
||||
"rareterm qdrant backend note",
|
||||
"frontend design note",
|
||||
],
|
||||
metadatas=[
|
||||
{"wing": "project", "room": "backend", "rank": 1},
|
||||
{"wing": "project", "room": "backend", "rank": 3},
|
||||
{"wing": "project", "room": "frontend", "rank": 2},
|
||||
],
|
||||
embeddings=[[1, 0], [0.9, 0.1], [0, 1]],
|
||||
)
|
||||
|
||||
assert QdrantBackend.detect(str(tmp_path))
|
||||
assert os.path.isfile(tmp_path / "qdrant_backend.json")
|
||||
assert col.count() == 3
|
||||
|
||||
result = col.query(
|
||||
query_embeddings=[[1, 0]],
|
||||
n_results=3,
|
||||
where={"rank": {"$gte": 2}},
|
||||
include=["documents", "metadatas", "distances", "embeddings"],
|
||||
)
|
||||
assert result.ids == [["b", "c"]]
|
||||
assert result.documents[0][0] == "rareterm qdrant backend note"
|
||||
assert result.embeddings[0][0] == pytest.approx([0.9, 0.1])
|
||||
|
||||
hits = col.lexical_search(query="rareterm backend", n_results=2, where={"wing": "project"}).hits
|
||||
assert [hit.id for hit in hits] == ["b", "a"]
|
||||
assert fake_qdrant.instances[0].created_indexes[0][1:] == ("document", "text")
|
||||
|
||||
backend.close_palace(str(tmp_path))
|
||||
with pytest.raises(Exception):
|
||||
col.count()
|
||||
|
||||
|
||||
def test_qdrant_marker_not_written_when_first_write_fails(tmp_path, fake_qdrant, monkeypatch):
|
||||
_backend, col = _collection(tmp_path)
|
||||
fake_client = fake_qdrant.instances[0]
|
||||
|
||||
def fail_upsert(*_args, **_kwargs):
|
||||
raise RuntimeError("qdrant unavailable")
|
||||
|
||||
monkeypatch.setattr(fake_client, "upsert_points", fail_upsert)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
col.upsert(ids=["a"], documents=["one"], metadatas=[{}], embeddings=[[1, 0]])
|
||||
|
||||
assert not os.path.isfile(tmp_path / "qdrant_backend.json")
|
||||
|
||||
|
||||
def test_qdrant_upsert_update_delete_get_order_and_multi_collection(tmp_path, fake_qdrant):
|
||||
backend, drawers = _collection(tmp_path, "drawers")
|
||||
palace = PalaceRef(id=str(tmp_path), local_path=str(tmp_path))
|
||||
closets = backend.get_collection(palace=palace, collection_name="closets", create=True)
|
||||
|
||||
drawers.upsert(
|
||||
ids=["one", "two"],
|
||||
documents=["first document", "second document"],
|
||||
metadatas=[{"wing": "a"}, {"wing": "b"}],
|
||||
embeddings=[[1, 0], [0, 1]],
|
||||
)
|
||||
closets.upsert(
|
||||
ids=["one"],
|
||||
documents=["closet document"],
|
||||
metadatas=[{"wing": "closet"}],
|
||||
embeddings=[[0.5, 0.5]],
|
||||
)
|
||||
|
||||
got = drawers.get(ids=["two", "one", "two"], include=["documents", "metadatas"])
|
||||
assert got.ids == ["two", "one", "two"]
|
||||
assert got.documents == ["second document", "first document", "second document"]
|
||||
|
||||
drawers.update(ids=["one"], metadatas=[{"room": "updated"}])
|
||||
assert drawers.get(ids=["one"]).metadatas == [{"wing": "a", "room": "updated"}]
|
||||
|
||||
drawers.delete(where={"wing": "b"})
|
||||
assert drawers.get().ids == ["one"]
|
||||
assert closets.get().ids == ["one"]
|
||||
|
||||
|
||||
def test_qdrant_complex_filters_use_exact_local_fallback(tmp_path, fake_qdrant):
|
||||
_backend, col = _collection(tmp_path)
|
||||
col.upsert(
|
||||
ids=["a", "b", "c"],
|
||||
documents=[
|
||||
"needle exact substring",
|
||||
"needle other wing",
|
||||
"boring filler",
|
||||
],
|
||||
metadatas=[
|
||||
{"wing": "target", "room": "backend", "tag": "alpha-beta"},
|
||||
{"wing": "other", "room": "backend", "tag": "beta"},
|
||||
{"wing": "target", "room": "front", "tag": "gamma"},
|
||||
],
|
||||
embeddings=[[1, 0], [0.8, 0.2], [0, 1]],
|
||||
)
|
||||
fake_client = fake_qdrant.instances[0]
|
||||
|
||||
result = col.query(
|
||||
query_embeddings=[[1, 0]],
|
||||
n_results=5,
|
||||
where={"$or": [{"wing": "target"}, {"tag": {"$contains": "alpha"}}]},
|
||||
where_document={"$contains": "needle"},
|
||||
)
|
||||
|
||||
assert result.ids == [["a"]]
|
||||
assert fake_client.query_calls == []
|
||||
|
||||
|
||||
def test_qdrant_dimension_mismatch(tmp_path, fake_qdrant):
|
||||
_backend, col = _collection(tmp_path)
|
||||
col.upsert(ids=["a"], documents=["one"], metadatas=[{}], embeddings=[[1, 0]])
|
||||
|
||||
with pytest.raises(DimensionMismatchError):
|
||||
col.upsert(ids=["b"], documents=["two"], metadatas=[{}], embeddings=[[1, 0, 0]])
|
||||
|
||||
|
||||
def test_qdrant_add_rejects_duplicate_ids_in_same_batch(tmp_path, fake_qdrant):
|
||||
_backend, col = _collection(tmp_path)
|
||||
|
||||
with pytest.raises(ValueError, match="unique"):
|
||||
col.add(
|
||||
ids=["dup", "dup"],
|
||||
documents=["first", "second"],
|
||||
metadatas=[{}, {}],
|
||||
embeddings=[[1, 0], [0, 1]],
|
||||
)
|
||||
|
||||
assert not os.path.isfile(tmp_path / "qdrant_backend.json")
|
||||
|
||||
|
||||
def test_qdrant_marker_participates_in_backend_mismatch(tmp_path, monkeypatch, fake_qdrant):
|
||||
from mempalace.palace import resolve_backend_name
|
||||
|
||||
backend, col = _collection(tmp_path)
|
||||
col.upsert(ids=["a"], documents=["one"], metadatas=[{}], embeddings=[[1, 0]])
|
||||
backend.close()
|
||||
(tmp_path / "chroma.sqlite3").write_bytes(b"")
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "chroma")
|
||||
|
||||
with pytest.raises(BackendMismatchError):
|
||||
resolve_backend_name(str(tmp_path))
|
||||
|
||||
|
||||
def test_qdrant_marker_rejects_remote_target_change(tmp_path, monkeypatch, fake_qdrant):
|
||||
backend, col = _collection(tmp_path)
|
||||
palace = PalaceRef(id=str(tmp_path), local_path=str(tmp_path))
|
||||
col.upsert(ids=["a"], documents=["one"], metadatas=[{}], embeddings=[[1, 0]])
|
||||
|
||||
monkeypatch.setenv("MEMPALACE_QDRANT_URL", "http://other-qdrant.example:6333")
|
||||
|
||||
with pytest.raises(BackendMismatchError, match="remote target"):
|
||||
backend.get_collection(palace=palace, collection_name="drawers", create=False)
|
||||
|
||||
|
||||
def test_qdrant_namespace_does_not_mix_palaces(tmp_path, fake_qdrant):
|
||||
backend = QdrantBackend()
|
||||
palace_a_path = tmp_path / "a"
|
||||
palace_b_path = tmp_path / "b"
|
||||
palace_a = PalaceRef(id=str(palace_a_path), local_path=str(palace_a_path), namespace="shared")
|
||||
palace_b = PalaceRef(id=str(palace_b_path), local_path=str(palace_b_path), namespace="shared")
|
||||
|
||||
col_a = backend.get_collection(palace=palace_a, collection_name="drawers", create=True)
|
||||
col_b = backend.get_collection(palace=palace_b, collection_name="drawers", create=True)
|
||||
col_a.upsert(ids=["same"], documents=["palace a"], metadatas=[{}], embeddings=[[1, 0]])
|
||||
col_b.upsert(ids=["same"], documents=["palace b"], metadatas=[{}], embeddings=[[1, 0]])
|
||||
|
||||
assert col_a.get(ids=["same"]).documents == ["palace a"]
|
||||
assert col_b.get(ids=["same"]).documents == ["palace b"]
|
||||
assert col_a._remote_collection != col_b._remote_collection
|
||||
|
||||
|
||||
def test_qdrant_missing_remote_after_marker_is_unhealthy(tmp_path, fake_qdrant):
|
||||
_backend, col = _collection(tmp_path)
|
||||
col.upsert(ids=["a"], documents=["one"], metadatas=[{}], embeddings=[[1, 0]])
|
||||
fake_client = fake_qdrant.instances[0]
|
||||
fake_client.delete_collection(col._remote_collection)
|
||||
|
||||
assert col.health().ok is False
|
||||
with pytest.raises(CollectionNotInitializedError):
|
||||
col.count()
|
||||
|
||||
|
||||
def test_search_reports_backend_error_distinct_from_missing_palace(tmp_path, monkeypatch):
|
||||
from mempalace import searcher
|
||||
|
||||
def fail_open(*_args, **_kwargs):
|
||||
raise BackendError("qdrant unavailable")
|
||||
|
||||
monkeypatch.setattr(searcher, "get_collection", fail_open)
|
||||
|
||||
result = searcher.search_memories("needle", str(tmp_path))
|
||||
|
||||
assert result["error"] == "Backend error"
|
||||
assert "qdrant unavailable" in result["details"]
|
||||
|
||||
|
||||
def test_palace_wrapper_embeds_for_qdrant(tmp_path, monkeypatch, fake_qdrant):
|
||||
import mempalace.backends.embedding_wrapper as embedding_wrapper
|
||||
from mempalace import palace
|
||||
|
||||
monkeypatch.setattr(
|
||||
embedding_wrapper, "_embed_texts", lambda texts: [[1.0, 0.0] for _ in texts]
|
||||
)
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "qdrant")
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND", "qdrant")
|
||||
|
||||
col = palace.get_collection(str(tmp_path), "mempalace_drawers", create=True)
|
||||
col.add(documents=["wrapped qdrant document"], ids=["wrapped"], metadatas=[{"wing": "w"}])
|
||||
result = col.query(query_texts=["wrapped"], n_results=1)
|
||||
assert result.ids == [["wrapped"]]
|
||||
|
||||
|
||||
def test_qdrant_live_rest_roundtrip_when_enabled(tmp_path):
|
||||
live_url = os.environ.get("MEMPALACE_QDRANT_LIVE_URL")
|
||||
if not live_url:
|
||||
pytest.skip("set MEMPALACE_QDRANT_LIVE_URL to run live Qdrant REST test")
|
||||
|
||||
backend = QdrantBackend()
|
||||
namespace = f"live_{uuid.uuid4().hex}"
|
||||
palace = PalaceRef(id=str(tmp_path), local_path=str(tmp_path), namespace=namespace)
|
||||
col = backend.get_collection(
|
||||
palace=palace,
|
||||
collection_name="drawers",
|
||||
create=True,
|
||||
options={
|
||||
"url": live_url,
|
||||
"api_key": os.environ.get("MEMPALACE_QDRANT_LIVE_API_KEY"),
|
||||
},
|
||||
)
|
||||
try:
|
||||
col.upsert(
|
||||
ids=["live-a", "live-b"],
|
||||
documents=["rareterm live qdrant backend", "other live document"],
|
||||
metadatas=[{"wing": "live", "rank": 2}, {"wing": "other", "rank": 1}],
|
||||
embeddings=[[1.0, 0.0], [0.0, 1.0]],
|
||||
)
|
||||
assert QdrantBackend.detect(str(tmp_path))
|
||||
|
||||
result = col.query(
|
||||
query_embeddings=[[1.0, 0.0]],
|
||||
n_results=2,
|
||||
where={"wing": "live"},
|
||||
)
|
||||
assert result.ids == [["live-a"]]
|
||||
|
||||
hits = col.lexical_search(query="rareterm", n_results=1).hits
|
||||
assert hits and hits[0].id == "live-a"
|
||||
|
||||
col.delete(ids=["live-a"])
|
||||
assert col.get(ids=["live-a"]).ids == []
|
||||
finally:
|
||||
try:
|
||||
col._client.delete_collection(col._remote_collection)
|
||||
except Exception:
|
||||
pass
|
||||
backend.close()
|
||||
|
|
@ -0,0 +1,360 @@
|
|||
import math
|
||||
|
||||
import pytest
|
||||
|
||||
from mempalace.backends import (
|
||||
BackendMismatchError,
|
||||
DimensionMismatchError,
|
||||
PalaceRef,
|
||||
QueryResult,
|
||||
UnsupportedCapabilityError,
|
||||
available_backends,
|
||||
)
|
||||
from mempalace.backends.sqlite_exact import SQLiteExactBackend
|
||||
|
||||
|
||||
def _collection(tmp_path, name="mempalace_drawers", create=True):
|
||||
backend = SQLiteExactBackend()
|
||||
palace = PalaceRef(id=str(tmp_path), local_path=str(tmp_path))
|
||||
return backend, backend.get_collection(palace=palace, collection_name=name, create=create)
|
||||
|
||||
|
||||
def test_registry_exposes_sqlite_exact():
|
||||
assert "sqlite_exact" in available_backends()
|
||||
|
||||
|
||||
def test_sqlite_exact_add_query_filters_and_persistence(tmp_path):
|
||||
backend, col = _collection(tmp_path)
|
||||
col.add(
|
||||
ids=["a", "b", "c"],
|
||||
documents=[
|
||||
"alpha vector memory",
|
||||
"beta sqlite exact memory",
|
||||
"gamma filtered memory",
|
||||
],
|
||||
metadatas=[
|
||||
{"wing": "alpha", "room": "notes", "chunk_index": 0, "tags": "core,vector"},
|
||||
{"wing": "alpha", "room": "notes", "chunk_index": 1, "tags": "sqlite,exact"},
|
||||
{"wing": "gamma", "room": "archive", "chunk_index": 2, "tags": "old"},
|
||||
],
|
||||
embeddings=[[1.0, 0.0], [0.0, 1.0], [0.2, 0.8]],
|
||||
)
|
||||
|
||||
ranked = col.query(query_embeddings=[[1.0, 0.0]], n_results=3)
|
||||
assert ranked.ids[0] == ["a", "c", "b"]
|
||||
assert ranked.distances[0][0] == pytest.approx(0.0)
|
||||
|
||||
filtered = col.get(
|
||||
where={
|
||||
"$and": [
|
||||
{"wing": "alpha"},
|
||||
{"chunk_index": {"$gte": 1}},
|
||||
{"tags": {"$contains": "sqlite"}},
|
||||
]
|
||||
},
|
||||
include=["documents", "metadatas", "embeddings"],
|
||||
)
|
||||
assert filtered.ids == ["b"]
|
||||
assert filtered.documents == ["beta sqlite exact memory"]
|
||||
assert filtered.embeddings == [[0.0, 1.0]]
|
||||
|
||||
col.update(ids=["b"], metadatas=[{"room": "lab"}])
|
||||
assert col.get(ids=["b"]).metadatas[0]["room"] == "lab"
|
||||
|
||||
backend.close_palace(str(tmp_path))
|
||||
reopened = backend.get_collection(
|
||||
palace=PalaceRef(id=str(tmp_path), local_path=str(tmp_path)),
|
||||
collection_name="mempalace_drawers",
|
||||
create=False,
|
||||
)
|
||||
assert reopened.count() == 3
|
||||
assert reopened.get(ids=["a"]).documents == ["alpha vector memory"]
|
||||
|
||||
|
||||
def test_sqlite_exact_write_failure_rolls_back_whole_batch(tmp_path):
|
||||
_backend, col = _collection(tmp_path)
|
||||
|
||||
with pytest.raises(Exception):
|
||||
col.add(
|
||||
ids=["dup", "dup"],
|
||||
documents=["first write", "duplicate write"],
|
||||
metadatas=[{}, {}],
|
||||
embeddings=[[1.0, 0.0], [0.0, 1.0]],
|
||||
)
|
||||
|
||||
assert col.count() == 0
|
||||
|
||||
|
||||
def test_sqlite_exact_enforces_collection_dimension(tmp_path):
|
||||
_backend, col = _collection(tmp_path)
|
||||
col.add(ids=["a"], documents=["two dims"], metadatas=[{}], embeddings=[[1.0, 0.0]])
|
||||
|
||||
with pytest.raises(DimensionMismatchError):
|
||||
col.add(ids=["b"], documents=["three dims"], metadatas=[{}], embeddings=[[1.0, 0.0, 0.0]])
|
||||
with pytest.raises(DimensionMismatchError):
|
||||
col.upsert(
|
||||
ids=["b"], documents=["three dims"], metadatas=[{}], embeddings=[[1.0, 0.0, 0.0]]
|
||||
)
|
||||
with pytest.raises(DimensionMismatchError):
|
||||
col.update(ids=["a"], embeddings=[[1.0, 0.0, 0.0]])
|
||||
with pytest.raises(DimensionMismatchError):
|
||||
col.query(query_embeddings=[[1.0, 0.0, 0.0]], n_results=1)
|
||||
|
||||
assert col.count() == 1
|
||||
assert col.get(ids=["a"]).documents == ["two dims"]
|
||||
|
||||
|
||||
def test_sqlite_exact_get_preserves_requested_id_order_and_duplicates(tmp_path):
|
||||
_backend, col = _collection(tmp_path)
|
||||
col.add(
|
||||
ids=["a", "b"],
|
||||
documents=["doc a", "doc b"],
|
||||
metadatas=[{}, {}],
|
||||
embeddings=[[1, 0], [0, 1]],
|
||||
)
|
||||
|
||||
result = col.get(ids=["b", "a", "b"], include=["documents"])
|
||||
|
||||
assert result.ids == ["b", "a", "b"]
|
||||
assert result.documents == ["doc b", "doc a", "doc b"]
|
||||
|
||||
|
||||
def test_sqlite_exact_upsert_delete_and_multi_collection_isolation(tmp_path):
|
||||
backend, drawers = _collection(tmp_path, "drawers")
|
||||
palace = PalaceRef(id=str(tmp_path), local_path=str(tmp_path))
|
||||
closets = backend.get_collection(palace=palace, collection_name="closets", create=True)
|
||||
|
||||
drawers.upsert(
|
||||
ids=["same"], documents=["drawer one"], metadatas=[{"kind": "drawer"}], embeddings=[[1, 0]]
|
||||
)
|
||||
closets.upsert(
|
||||
ids=["same"], documents=["closet one"], metadatas=[{"kind": "closet"}], embeddings=[[0, 1]]
|
||||
)
|
||||
drawers.upsert(
|
||||
ids=["same"],
|
||||
documents=["drawer replaced"],
|
||||
metadatas=[{"kind": "drawer", "version": 2}],
|
||||
embeddings=[[1, 0]],
|
||||
)
|
||||
|
||||
assert drawers.count() == 1
|
||||
assert closets.count() == 1
|
||||
assert drawers.get(ids=["same"]).documents == ["drawer replaced"]
|
||||
assert closets.get(ids=["same"]).documents == ["closet one"]
|
||||
|
||||
drawers.delete(where={"version": {"$in": [2, 3]}})
|
||||
assert drawers.count() == 0
|
||||
assert closets.count() == 1
|
||||
|
||||
|
||||
def test_sqlite_exact_lexical_search_and_python_fallback(tmp_path, monkeypatch):
|
||||
_backend, col = _collection(tmp_path)
|
||||
col.add(
|
||||
ids=["a", "b", "c"],
|
||||
documents=[
|
||||
"ordinary project note",
|
||||
"rareterm rareterm sqlite exact note",
|
||||
"rareterm unrelated archive",
|
||||
],
|
||||
metadatas=[
|
||||
{"wing": "w", "room": "a"},
|
||||
{"wing": "w", "room": "b"},
|
||||
{"wing": "old", "room": "b"},
|
||||
],
|
||||
embeddings=[[1, 0], [0, 1], [0.5, 0.5]],
|
||||
)
|
||||
|
||||
hits = col.lexical_search(query="rareterm sqlite", n_results=2, where={"wing": "w"}).hits
|
||||
assert [hit.id for hit in hits] == ["b"]
|
||||
|
||||
monkeypatch.setattr(col, "_fts_available", lambda _cur: False)
|
||||
fallback_hits = col.lexical_search(query="rareterm sqlite", n_results=2).hits
|
||||
assert fallback_hits[0].id == "b"
|
||||
|
||||
|
||||
def test_sqlite_exact_lexical_search_filters_after_full_fts_window(tmp_path):
|
||||
_backend, col = _collection(tmp_path)
|
||||
ids = [f"old-{i}" for i in range(12)] + ["target"]
|
||||
col.add(
|
||||
ids=ids,
|
||||
documents=["needle shared lexical note" for _ in ids],
|
||||
metadatas=[{"wing": "old"} for _ in range(12)] + [{"wing": "target"}],
|
||||
embeddings=[[1.0, 0.0] for _ in ids],
|
||||
)
|
||||
|
||||
hits = col.lexical_search(query="needle", n_results=1, where={"wing": "target"}).hits
|
||||
|
||||
assert [hit.id for hit in hits] == ["target"]
|
||||
|
||||
|
||||
def test_sqlite_exact_logical_filters_evaluate_sibling_predicates(tmp_path):
|
||||
_backend, col = _collection(tmp_path)
|
||||
col.add(
|
||||
ids=["a", "b"],
|
||||
documents=["alpha document", "beta document"],
|
||||
metadatas=[
|
||||
{"wing": "w", "room": "wrong", "kind": "note"},
|
||||
{"wing": "w", "room": "right", "kind": "note"},
|
||||
],
|
||||
embeddings=[[1, 0], [0, 1]],
|
||||
)
|
||||
|
||||
result = col.get(where={"$and": [{"wing": "w"}], "room": "right"})
|
||||
|
||||
assert result.ids == ["b"]
|
||||
|
||||
|
||||
def test_sqlite_exact_close_palace_marks_existing_collections_closed(tmp_path):
|
||||
backend, col = _collection(tmp_path)
|
||||
palace = PalaceRef(id=str(tmp_path), local_path=str(tmp_path))
|
||||
col.add(ids=["a"], documents=["doc"], metadatas=[{}], embeddings=[[1, 0]])
|
||||
|
||||
backend.close_palace(palace)
|
||||
|
||||
assert not col.health().ok
|
||||
with pytest.raises(Exception):
|
||||
col.count()
|
||||
|
||||
|
||||
def test_palace_wrapper_embeds_for_sqlite_exact(tmp_path, monkeypatch):
|
||||
import mempalace.backends.embedding_wrapper as embedding_wrapper
|
||||
from mempalace.palace import get_collection
|
||||
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "sqlite_exact")
|
||||
monkeypatch.setattr(
|
||||
embedding_wrapper,
|
||||
"_embed_texts",
|
||||
lambda texts: [[float(len(text)), 1.0] for text in texts],
|
||||
)
|
||||
|
||||
col = get_collection(str(tmp_path), create=True)
|
||||
col.add(ids=["a"], documents=["abcd"], metadatas=[{"wing": "w"}])
|
||||
|
||||
result = col.query(query_texts=["abcd"], n_results=1)
|
||||
assert result.ids == [["a"]]
|
||||
|
||||
|
||||
def test_backend_mismatch_protection(tmp_path, monkeypatch):
|
||||
from mempalace.palace import get_collection
|
||||
|
||||
(tmp_path / "chroma.sqlite3").write_bytes(b"")
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "sqlite_exact")
|
||||
|
||||
with pytest.raises(BackendMismatchError):
|
||||
get_collection(str(tmp_path), create=True)
|
||||
|
||||
|
||||
def test_mixed_backend_artifacts_are_rejected_even_when_chroma_selected(tmp_path, monkeypatch):
|
||||
from mempalace.palace import resolve_backend_name
|
||||
|
||||
(tmp_path / "chroma.sqlite3").write_bytes(b"")
|
||||
(tmp_path / "sqlite_exact.sqlite3").write_bytes(b"")
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "chroma")
|
||||
|
||||
with pytest.raises(BackendMismatchError):
|
||||
resolve_backend_name(str(tmp_path))
|
||||
|
||||
|
||||
def test_sqlite_exact_exact_ranking_uses_cosine(tmp_path):
|
||||
_backend, col = _collection(tmp_path)
|
||||
halfway = [0.5, math.sqrt(0.75)]
|
||||
col.add(
|
||||
ids=["half", "orthogonal", "same"],
|
||||
documents=["half", "orthogonal", "same"],
|
||||
metadatas=[{}, {}, {}],
|
||||
embeddings=[halfway, [0.0, 1.0], [1.0, 0.0]],
|
||||
)
|
||||
|
||||
result = col.query(query_embeddings=[[1.0, 0.0]], n_results=3)
|
||||
assert result.ids[0] == ["same", "half", "orthogonal"]
|
||||
assert result.distances[0] == pytest.approx([0.0, 0.5, 1.0])
|
||||
|
||||
|
||||
def test_search_union_uses_sqlite_exact_lexical_search(tmp_path, monkeypatch):
|
||||
import mempalace.backends.embedding_wrapper as embedding_wrapper
|
||||
from mempalace.palace import get_collection
|
||||
from mempalace.searcher import search_memories
|
||||
|
||||
def fake_embed(texts):
|
||||
vectors = []
|
||||
for text in texts:
|
||||
if text == "rareterm":
|
||||
vectors.append([1.0, 0.0])
|
||||
elif "rareterm" in text:
|
||||
vectors.append([0.0, 1.0])
|
||||
else:
|
||||
vectors.append([0.5, math.sqrt(0.75)])
|
||||
return vectors
|
||||
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "sqlite_exact")
|
||||
monkeypatch.setattr(embedding_wrapper, "_embed_texts", fake_embed)
|
||||
|
||||
col = get_collection(str(tmp_path), create=True)
|
||||
col.add(
|
||||
ids=["d1", "d2", "d3", "rare"],
|
||||
documents=[
|
||||
"ordinary support note",
|
||||
"ordinary billing note",
|
||||
"ordinary project note",
|
||||
"rareterm rareterm rareterm policy note",
|
||||
],
|
||||
metadatas=[
|
||||
{"wing": "w", "room": "r", "source_file": "/tmp/d1.md", "chunk_index": 0},
|
||||
{"wing": "w", "room": "r", "source_file": "/tmp/d2.md", "chunk_index": 0},
|
||||
{"wing": "w", "room": "r", "source_file": "/tmp/d3.md", "chunk_index": 0},
|
||||
{"wing": "w", "room": "r", "source_file": "/tmp/rare.md", "chunk_index": 0},
|
||||
],
|
||||
)
|
||||
|
||||
result = search_memories(
|
||||
"rareterm",
|
||||
str(tmp_path),
|
||||
n_results=1,
|
||||
candidate_strategy="union",
|
||||
)
|
||||
|
||||
assert result["results"][0]["source_file"] == "rare.md"
|
||||
assert result["results"][0]["matched_via"] == "bm25_backend"
|
||||
|
||||
|
||||
def test_search_union_reports_unsupported_lexical_capability(monkeypatch, tmp_path):
|
||||
import mempalace.searcher as searcher
|
||||
|
||||
class NoLexicalCollection:
|
||||
def query(self, **_kwargs):
|
||||
return QueryResult(
|
||||
ids=[["a"]],
|
||||
documents=[["ordinary note"]],
|
||||
metadatas=[[{"source_file": "/tmp/a.md", "chunk_index": 0}]],
|
||||
distances=[[0.5]],
|
||||
)
|
||||
|
||||
def lexical_search(self, **_kwargs):
|
||||
raise UnsupportedCapabilityError("no lexical support")
|
||||
|
||||
monkeypatch.setattr(searcher, "get_collection", lambda *_args, **_kwargs: NoLexicalCollection())
|
||||
monkeypatch.setattr(
|
||||
searcher,
|
||||
"get_closets_collection",
|
||||
lambda *_args, **_kwargs: (_ for _ in ()).throw(RuntimeError("no closets")),
|
||||
)
|
||||
|
||||
result = searcher.search_memories(
|
||||
"anything",
|
||||
str(tmp_path),
|
||||
n_results=1,
|
||||
candidate_strategy="union",
|
||||
)
|
||||
|
||||
assert result["unsupported_capability"] == "supports_lexical_search"
|
||||
|
||||
|
||||
def test_search_vector_disabled_fallback_is_chroma_only(tmp_path, monkeypatch):
|
||||
from mempalace.searcher import search_memories
|
||||
|
||||
monkeypatch.setenv("MEMPALACE_BACKEND_EXPLICIT", "sqlite_exact")
|
||||
|
||||
result = search_memories("anything", str(tmp_path), vector_disabled=True)
|
||||
|
||||
assert result["unsupported_capability"] == "chroma_hnsw_fallback"
|
||||
assert result["backend"] == "sqlite_exact"
|
||||
42
uv.lock
42
uv.lock
|
|
@ -2036,7 +2036,7 @@ requires-dist = [
|
|||
{ name = "pytest-cov", marker = "extra == 'dev'", specifier = ">=4.0" },
|
||||
{ name = "python-dateutil", specifier = ">=2.8" },
|
||||
{ name = "pyyaml", specifier = ">=6.0,<7" },
|
||||
{ name = "ruff", marker = "extra == 'dev'", specifier = "==0.15.14" },
|
||||
{ name = "ruff", marker = "extra == 'dev'", specifier = "==0.15.15" },
|
||||
{ name = "striprtf", marker = "extra == 'extract'", specifier = ">=0.0.27" },
|
||||
{ name = "tokenizers", specifier = ">=0.15" },
|
||||
{ name = "tomli", marker = "python_full_version < '3.11'", specifier = ">=2.0.0" },
|
||||
|
|
@ -2051,7 +2051,7 @@ dev = [
|
|||
{ name = "psutil", specifier = ">=5.9" },
|
||||
{ name = "pytest", specifier = ">=7.0" },
|
||||
{ name = "pytest-cov", specifier = ">=4.0" },
|
||||
{ name = "ruff", specifier = "==0.15.14" },
|
||||
{ name = "ruff", specifier = "==0.15.15" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -4821,27 +4821,27 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "ruff"
|
||||
version = "0.15.14"
|
||||
version = "0.15.15"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/dc/8a/8bce2894573e9dae6ff4d77fe34ad727d79b9e6238ad288c5638990d90f6/ruff-0.15.14.tar.gz", hash = "sha256:48e866b165be4a9bdbf310f7d3c9a07edef2fe8cd63ffeb4e00bb590506ebf9f", size = 4700910, upload-time = "2026-05-21T14:34:55.177Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/84/6f/a76f7d96e5c962f5b69cee865e49c15c1116897c01990faa8a57edb62e7f/ruff-0.15.15.tar.gz", hash = "sha256:b8dff018130b46d8e5bf0f926ef6b60cf871d6d5ae45fc9334e09632daa741d6", size = 4706985, upload-time = "2026-05-28T14:16:57.784Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/b9/c8/74a92c6ff9fcfb4f1f947126d3ebee8389276e161ecc85de5bda7cda51bd/ruff-0.15.14-py3-none-linux_armv6l.whl", hash = "sha256:8dd2db9416e487c8d4b01fa7056bb02c4d05969d4f8d17a08c229c2f4ff3c108", size = 10739177, upload-time = "2026-05-21T14:34:37.332Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/45/91/254a35c20acc38a7223c9d2d594af12e794432464f2cdeb52af1dc4a892d/ruff-0.15.14-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:be4ff55af755bd71a00ab3dc6bd7ffc467bd76e0df6881e286c2e3d23e8fb43b", size = 11144969, upload-time = "2026-05-21T14:34:43.978Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/56/9e/d13e40f83b8d0a94430e6778ce1d94a43b38cf2efe63278bdd2b4c65abbf/ruff-0.15.14-py3-none-macosx_11_0_arm64.whl", hash = "sha256:48d5909d7d06276ce7dde6d32bfa4b0d4cb2651145cd8ee4b440722cbc77832f", size = 10478207, upload-time = "2026-05-21T14:34:48.378Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/8d/f1/b15a7839fa4f332f8acec78e20564f26bb2d866e3d21710b877fd0263000/ruff-0.15.14-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:ca8cbfa94c4f90984a67561978602746d4cd27103568f745fa90eee3f0d4107d", size = 10818459, upload-time = "2026-05-21T14:34:22.318Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/45/33/53d651177f84f94b400a0e27f8824eeada3dddc9d5ee8aeb048f4352a520/ruff-0.15.14-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:9a6bbc0333f1ab053423bcbf6226477d266ca7cec7738c4c8e3f55647803f3c4", size = 10541800, upload-time = "2026-05-21T14:34:20.209Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/b8/a6/868f87e0bf9786ed24b5d0d0ad8676b8a94fd1912f42cddf9cfc7857818a/ruff-0.15.14-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:8a24a4f7605d7003a6674d4387651effd939dead3fddd0f36561eb77a9a2e542", size = 11342149, upload-time = "2026-05-21T14:34:46.365Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/a7/8b/38cd5c19faffdcc05a408d2b78edccc69492ab9720eadb49ea15ef80d768/ruff-0.15.14-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:049b5326e53ed80978f2fc041a280603f69dd6b0c95464342a2bb4572d9d9e2f", size = 12212563, upload-time = "2026-05-21T14:34:28.579Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/3e/4d/a3c5b874a556d5731e3e657aaf04311bb76f0a5c3ec220ed43051be6b64b/ruff-0.15.14-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:d4ed42e6696c8dfa5f06728e6441993901f548eb92d73bc472cb5a38d1395fbf", size = 11493299, upload-time = "2026-05-21T14:34:41.836Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/1e/c0/56472c251d09858a53e51efbd485b09e1995d8731668b76d52e5dd6ee0f1/ruff-0.15.14-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:715c543cf450c4888251f91c52f1942a800541d9bddd7ac060aa4e6b77ae7cba", size = 11455931, upload-time = "2026-05-21T14:34:57.276Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/2c/4a/e2e7b4d8dbf233d4eace59c75bc3435fa6d8bd3bae82d351d4e4300c0fd1/ruff-0.15.14-py3-none-manylinux_2_31_riscv64.whl", hash = "sha256:72ebab6013ec887d439d8b7593737a0a4ffb06d45d209d4e4bf2e92813082d3f", size = 11400794, upload-time = "2026-05-21T14:34:39.773Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/97/c7/83c0539fe34c3e09136204d1e75d6052492364e0b3cb05e9465423f567d7/ruff-0.15.14-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:49072d36abdbe97a8dd7f480afe9c675699c0c495d4c84076e2c1203c4550581", size = 10804759, upload-time = "2026-05-21T14:34:31.045Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/86/a6/18f2bfc095a2ab4a78745644e428205532ce6653a5d0fa8501572891534d/ruff-0.15.14-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:958522aee105068640c2c2ceae08f413ae44d922f52a1374ac13d6a96032fc93", size = 10539517, upload-time = "2026-05-21T14:34:53.064Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/54/3a/5a8b3b69c654d4e4bf1d246ac5b49cbcdac6eaab6905925f8915f31e3b80/ruff-0.15.14-py3-none-musllinux_1_2_i686.whl", hash = "sha256:f3707da619a143a2e8830e2abab8224478d69ace2d28cb6c20543ae97c36bf61", size = 11065169, upload-time = "2026-05-21T14:34:24.484Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/ed/c5/8864e4e7925b836ea354b31d57641ec03830564e281a8b6f061f8c3e0ec1/ruff-0.15.14-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:bb01d645694e3ec0102105d07ef2d53703970407d59c04e59d3ba0b7a1d53553", size = 11560214, upload-time = "2026-05-21T14:34:50.975Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/36/38/012bf76752e1f89ed50b77b99532d90f3a3e287bc7918e1fc0948ac866ac/ruff-0.15.14-py3-none-win32.whl", hash = "sha256:6d0c1ad2a0ab718d39b6d8fd2217981ce4d625cd96a720095f798fb47d8b13e6", size = 10805548, upload-time = "2026-05-21T14:34:33.453Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/d1/b7/4ea2c170f10ad760fff2a5250beb18897719dc8b52b53a24cddbb9dd3f19/ruff-0.15.14-py3-none-win_amd64.whl", hash = "sha256:802342981e056db3851a7836e5b070f8f15f67d4a685ae2a6160939d364b2902", size = 11939523, upload-time = "2026-05-21T14:34:18.077Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/62/d5/bc97ff895ec35cf3925d4bd60f3b39d822f377a446906ec9bcc87405e59b/ruff-0.15.14-py3-none-win_arm64.whl", hash = "sha256:ff47b90a9ef6a40c9e2f3b479c1fb78531adf055b94c1eba0a7ba04b31951826", size = 11208607, upload-time = "2026-05-21T14:34:26.525Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/fa/9d/3a45c05b8ab04b4705989de70a79008e27c8003296a0feaee9edc18dd7e9/ruff-0.15.15-py3-none-linux_armv6l.whl", hash = "sha256:cf93e5388f412e1b108b1f8b34a6e036b70fe8aff89393befad96fe48670311b", size = 10710652, upload-time = "2026-05-28T14:16:06.701Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/05/66/da974431624bf3b49f6ee1f9543c02d929ff1cba78b0d5a79c38cf21f744/ruff-0.15.15-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:ac5a646d1f6a7dadd5d50842dae2c1f9862ac887ef5d1b1375e02def791fde6e", size = 11096615, upload-time = "2026-05-28T14:16:23.313Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/8c/09/7443452e5d290230a712103f2fdceeef7184f3ec99a2bd01c8be78aaceb5/ruff-0.15.15-py3-none-macosx_11_0_arm64.whl", hash = "sha256:77d955a431430c66f72dd94e379ad38a16daea3d25094872ac4edf9e797be530", size = 10436683, upload-time = "2026-05-28T14:16:40.974Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/53/01/d330c26a57fa4f3943a14424904027428315b700fe4d14a84bb123a649e5/ruff-0.15.15-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7614ee79c69788cf6cedd568069ade9cecc22a1ad20494efe8d0c9ebb4b622d4", size = 10769064, upload-time = "2026-05-28T14:16:28.905Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/1d/85/cc8770f8bdff541b1da8392d1634141fe4a0e3f4ee596605959b7906c27f/ruff-0.15.15-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:3cdb1679e06a1f6b47bc384714ae96f6e2fb65ca441eb78c43d2ca554176ce1f", size = 10511987, upload-time = "2026-05-28T14:16:43.732Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/7c/29/8c190c1472b63013583ba391f3342036e02010544c1270455ed8e519bdf3/ruff-0.15.15-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:2728b93d7b23a603ea2c0ac6eb73d760bd38ec9de35f35fb41e18f7a3fee7622", size = 11275100, upload-time = "2026-05-28T14:16:55.244Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/9f/6b/7e145ce2cc8e63d6834eca03d83a0e18d121def5c69f91b4cf4011ed4879/ruff-0.15.15-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:be582fcc0db438902c7792b08d6ddf6c9b9e21addaa10092c2c741cfb09e5a45", size = 12176903, upload-time = "2026-05-28T14:16:14.368Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/80/a3/d5974637f68e451f7fadf015cf3101d1cd7d8ba5027cffe0b9e3826ebe6b/ruff-0.15.15-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:7aa77465b8ecaf1a27bea098d696f7fed5e1eccbd10b321b682d6de586ae5627", size = 11404550, upload-time = "2026-05-28T14:16:20.138Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/fe/1c/e6e5e568f22be4fb05d6244234aba384c06b451252453b821e1a529263cf/ruff-0.15.15-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:48decfa11d740de4889de623be1463308346312f2409a56e24aa280c86162dc4", size = 11382027, upload-time = "2026-05-28T14:16:46.615Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/1d/01/170921b49fcd2e8858825593f91cf7146c3e40a5c3e6df763e4bb0484dde/ruff-0.15.15-py3-none-manylinux_2_31_riscv64.whl", hash = "sha256:a5015088452ca0081387063649ec67f06d3d1d6b8b936a1f836b5e9657ecd48c", size = 11366041, upload-time = "2026-05-28T14:16:26.247Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/87/54/a7bad711d7de93254e15e06a4c375b89a03d18de45d3e5dcc86a4472fb1a/ruff-0.15.15-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:f5294aab6356c81600fcdea3a62bb1b924dfd5e91767c12318d3f68f86af57cd", size = 10741795, upload-time = "2026-05-28T14:16:17.11Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/c9/31/38c075963668f8b41c6914ee0f6f318727fbe30ab9145cb29e6df464c5fa/ruff-0.15.15-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:db5bd4d802415cca656dc1616070b725952d6ae95eb5d4831e49fbd94a38f75f", size = 10511117, upload-time = "2026-05-28T14:16:31.767Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/9d/96/6ff689e1f7e375d1d97075eca022f74c2bab59554a432fe4d2e6f091986a/ruff-0.15.15-py3-none-musllinux_1_2_i686.whl", hash = "sha256:587a6278ed42059191c1a466e490bd7930fb50bd2e255398bc29616c895a61cb", size = 10994867, upload-time = "2026-05-28T14:16:35.149Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/c3/c2/5dce0ab9f92a8d534fa62b9bf9caca3eddb8c1a81b616f5e195ada4f0d6e/ruff-0.15.15-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:df0c1c084f5f4be9812f61518a45c440d3c30d69ce4bf6c5270e66d38338f02a", size = 11482101, upload-time = "2026-05-28T14:16:49.598Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/b1/c0/1003b60edd697c649faf61f1a34094b1abb38fb3d1181e3f895781250a08/ruff-0.15.15-py3-none-win32.whl", hash = "sha256:29428ea79694afbe756d45fd59b36f22b6b020dc0443cf7de0173046236964b9", size = 10716774, upload-time = "2026-05-28T14:16:52.337Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/02/a8/1269eddd6945a06c23f055ef7848886e37cf9d6a8bebb386a3115f01470c/ruff-0.15.15-py3-none-win_amd64.whl", hash = "sha256:8df0323902e15e24bc4bf246da830573d3cf3352bd0b9a164eab335d111ff4a4", size = 11868463, upload-time = "2026-05-28T14:16:11.333Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/4e/b2/920464c907b191e37469d477a1aa8bc048b8f36c4c1610dfa4ab87b39e18/ruff-0.15.15-py3-none-win_arm64.whl", hash = "sha256:3c8ceca6792f38196b8f589bc92eccd03eef286602da92e5dc05cc42ef6441b7", size = 11138498, upload-time = "2026-05-28T14:16:38.425Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
Loading…
Reference in New Issue