Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 36 additions & 6 deletions mempalace/backends/embedding_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Comment on lines +21 to +35

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

Currently, _as_list only handles str by wrapping it in a list. However, ChromaDB's OneOrMany shape also allows a single dictionary for metadatas (e.g., {"source": "web"}). If a single dictionary is passed, calling list(value) on it will extract only its keys (e.g., ["source"]), which leads to data loss and schema mismatch.

By extending _as_list to also check for dict, we can safely wrap single metadata dictionaries as well.

Suggested change
def _as_list(value):
"""Normalize ChromaDB's ``OneOrMany`` shape (``str`` | sequence) to a list.
A bare ``str`` must be *wrapped*, not iterated: ``list("abc")`` yields
``['a', 'b', 'c']``, which would embed per character and break length
alignment with ``ids``/``metadatas`` on explicit-vector backends
(pgvector, sqlite_exact). See PR #1706 review.
"""
if isinstance(value, str):
return [value]
return list(value)
def _as_list(value):
"""Normalize ChromaDB's ``OneOrMany`` shape (``str`` | ``dict`` | sequence) to a list.
A bare ``str`` or ``dict`` must be *wrapped*, not iterated: ``list("abc")`` yields
``['a', 'b', 'c']``, and ``list({"a": 1})`` yields ``['a']``, which would break
length alignment with ``ids``/``metadatas`` on explicit-vector backends
(pgvector, sqlite_exact). See PR #1706 review.
"""
if isinstance(value, (str, dict)):
return [value]
return list(value)



class EmbeddingCollection(BaseCollection):
"""Wrap a collection that requires explicit vectors.

Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -55,15 +80,15 @@ 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,
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_embeddings = _embed_texts(_as_list(query_texts))
query_texts = None
return self._inner.query(
query_texts=query_texts,
Expand Down Expand Up @@ -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,
Expand Down
126 changes: 126 additions & 0 deletions tests/test_embedding_wrapper.py
Original file line number Diff line number Diff line change
@@ -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
Comment on lines +108 to +113

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Let's add a unit test to verify that bare string ids and single dictionary metadatas are correctly wrapped and normalized before delegating to the inner backend.

Suggested change
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_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):
seen = _patch_embed(monkeypatch)
inner = _FakeInner()
ew.EmbeddingCollection(inner).add(
documents="hello world",
ids="d1",
metadatas={"source": "web"}
)
assert inner.calls["add"]["ids"] == ["d1"]
assert inner.calls["add"]["metadatas"] == [{"source": "web"}]



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
Loading