diff --git a/mempalace/backends/embedding_wrapper.py b/mempalace/backends/embedding_wrapper.py index f4a8ccda09..1441deb2de 100644 --- a/mempalace/backends/embedding_wrapper.py +++ b/mempalace/backends/embedding_wrapper.py @@ -18,6 +18,23 @@ def _embed_texts(texts: list[str]) -> list[list[float]]: return [list(v) for v in vectors] +def _as_list(value): + """Normalize ChromaDB's ``OneOrMany`` shape (``str`` | ``dict`` | sequence) to a list. + + A bare ``str`` (a document/id) or ``dict`` (a single metadata) must be + *wrapped*, not iterated: ``list("abc")`` yields ``['a', 'b', 'c']`` and + ``list({"k": 1})`` yields ``['k']`` — either desyncs embeddings/metadatas + from ``ids`` on explicit-vector backends (pgvector, sqlite_exact). A list is + returned unchanged (no copy); any other iterable is materialized once. + See PR #1706/#1707 review. + """ + if isinstance(value, (str, dict)): + return [value] + if isinstance(value, list): + return value + return list(value) + + class EmbeddingCollection(BaseCollection): """Wrap a collection that requires explicit vectors. @@ -33,8 +50,12 @@ def __getattr__(self, name): return getattr(self._inner, name) def add(self, *, documents, ids, metadatas=None, embeddings=None): + documents = _as_list(documents) + ids = _as_list(ids) + if metadatas is not None: + metadatas = _as_list(metadatas) if embeddings is None: - embeddings = _embed_texts(list(documents)) + embeddings = _embed_texts(documents) return self._inner.add( documents=documents, ids=ids, @@ -43,8 +64,12 @@ def add(self, *, documents, ids, metadatas=None, embeddings=None): ) def upsert(self, *, documents, ids, metadatas=None, embeddings=None): + documents = _as_list(documents) + ids = _as_list(ids) + if metadatas is not None: + metadatas = _as_list(metadatas) if embeddings is None: - embeddings = _embed_texts(list(documents)) + embeddings = _embed_texts(documents) return self._inner.upsert( documents=documents, ids=ids, @@ -55,7 +80,7 @@ def upsert(self, *, documents, ids, metadatas=None, embeddings=None): def query( self, *, - query_texts: Optional[list[str]] = None, + query_texts: Optional[list[str] | str] = None, query_embeddings: Optional[list[list[float]]] = None, n_results: int = 10, where: Optional[dict] = None, @@ -63,7 +88,7 @@ def query( include: Optional[list[str]] = None, ): if query_texts is not None and query_embeddings is None: - query_embeddings = _embed_texts(list(query_texts)) + query_embeddings = _embed_texts(_as_list(query_texts)) query_texts = None return self._inner.query( query_texts=query_texts, @@ -105,8 +130,13 @@ def lexical_search(self, *, query: str, n_results: int = 10, where: Optional[dic 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)) + ids = _as_list(ids) + if documents is not None: + documents = _as_list(documents) + if embeddings is None: + embeddings = _embed_texts(documents) + if metadatas is not None: + metadatas = _as_list(metadatas) return self._inner.update( ids=ids, documents=documents, diff --git a/tests/test_embedding_wrapper.py b/tests/test_embedding_wrapper.py new file mode 100644 index 0000000000..5d592d6afa --- /dev/null +++ b/tests/test_embedding_wrapper.py @@ -0,0 +1,126 @@ +"""EmbeddingCollection OneOrMany handling (PR #1706 review). + +A bare ``str`` passed as ``documents``/``query_texts`` (ChromaDB's OneOrMany +shape) must be wrapped, not iterated — otherwise ``list("abc")`` embeds per +character and breaks length alignment with ids/metadatas on explicit-vector +backends. +""" + +from mempalace.backends import embedding_wrapper as ew + + +class _FakeInner: + """Captures what the wrapper delegates to the backend.""" + + def __init__(self): + self.calls = {} + + def add(self, *, documents, ids, metadatas=None, embeddings=None): + self.calls["add"] = { + "documents": documents, + "ids": ids, + "metadatas": metadatas, + "embeddings": embeddings, + } + + def upsert(self, *, documents, ids, metadatas=None, embeddings=None): + self.calls["upsert"] = { + "documents": documents, + "ids": ids, + "metadatas": metadatas, + "embeddings": embeddings, + } + + def update(self, *, ids, documents=None, metadatas=None, embeddings=None): + self.calls["update"] = { + "documents": documents, + "ids": ids, + "metadatas": metadatas, + "embeddings": embeddings, + } + + def query(self, *, query_texts=None, query_embeddings=None, **_kw): + self.calls["query"] = {"query_texts": query_texts, "query_embeddings": query_embeddings} + from mempalace.backends.base import QueryResult + + return QueryResult.empty() + + +def _patch_embed(monkeypatch): + """Stub the embedder: one vector per input text, recording the inputs.""" + seen = {} + + def fake(texts): + seen["texts"] = texts + return [[0.0, 0.0] for _ in texts] + + monkeypatch.setattr(ew, "_embed_texts", fake) + return seen + + +def test_as_list_wraps_bare_string(): + assert ew._as_list("hello world") == ["hello world"] + assert ew._as_list({"k": 1}) == [{"k": 1}] # bare dict wrapped, not -> ["k"] + src = ["a", "b"] + assert ew._as_list(src) is src # list returned as-is (no copy) + assert ew._as_list(("a", "b")) == ["a", "b"] # other iterables materialized + + +def test_add_wraps_bare_string_document(monkeypatch): + seen = _patch_embed(monkeypatch) + inner = _FakeInner() + ew.EmbeddingCollection(inner).add(documents="hello world", ids=["d1"]) + # embedded as one whole document, not per character + assert seen["texts"] == ["hello world"] + # and the backend receives a list, length-aligned with ids + assert inner.calls["add"]["documents"] == ["hello world"] + assert len(inner.calls["add"]["embeddings"]) == 1 + + +def test_upsert_wraps_bare_string_document(monkeypatch): + seen = _patch_embed(monkeypatch) + inner = _FakeInner() + ew.EmbeddingCollection(inner).upsert(documents="solo", ids=["d1"]) + assert seen["texts"] == ["solo"] + assert inner.calls["upsert"]["documents"] == ["solo"] + assert len(inner.calls["upsert"]["embeddings"]) == 1 + + +def test_update_wraps_bare_string_document(monkeypatch): + seen = _patch_embed(monkeypatch) + inner = _FakeInner() + ew.EmbeddingCollection(inner).update(ids=["d1"], documents="changed") + assert seen["texts"] == ["changed"] + assert inner.calls["update"]["documents"] == ["changed"] + assert len(inner.calls["update"]["embeddings"]) == 1 + + +def test_query_wraps_bare_string(monkeypatch): + seen = _patch_embed(monkeypatch) + inner = _FakeInner() + ew.EmbeddingCollection(inner).query(query_texts="find me") + assert seen["texts"] == ["find me"] + # query_texts is consumed into a single query embedding + assert len(inner.calls["query"]["query_embeddings"]) == 1 + assert inner.calls["query"]["query_texts"] is None + + +def test_list_inputs_unaffected(monkeypatch): + seen = _patch_embed(monkeypatch) + inner = _FakeInner() + ew.EmbeddingCollection(inner).add(documents=["one", "two"], ids=["a", "b"]) + assert seen["texts"] == ["one", "two"] + assert len(inner.calls["add"]["embeddings"]) == 2 + + +def test_add_wraps_bare_string_ids_and_dict_metadatas(monkeypatch): + _patch_embed(monkeypatch) + inner = _FakeInner() + # a single id (str) and a single metadata (dict) are OneOrMany shapes too + ew.EmbeddingCollection(inner).add(documents="solo", ids="d1", metadatas={"src": "web"}) + call = inner.calls["add"] + assert call["ids"] == ["d1"] # not ['d', '1'] + assert call["metadatas"] == [{"src": "web"}] # not ['src'] + # documents / embeddings / ids / metadatas all length-aligned at 1 + assert call["documents"] == ["solo"] + assert len(call["embeddings"]) == 1