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
50 changes: 48 additions & 2 deletions chromadol/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,8 +50,15 @@ def _ids(self):
def __iter__(self):
return iter(self._ids)

#: The ``include`` that reads pass to ``collection.get``. ``None`` keeps
#: chromadb's default (documents and metadatas); a single-field store whose
#: field is not in that default (e.g. ``uris``) names it here.
get_include: tuple[str, ...] | None = None

def __getitem__(self, k: str) -> GetResult:
return self.collection.get(k)
if self.get_include is None:
return self.collection.get(k)
return self.collection.get(k, include=list(self.get_include))

def __len__(self):
return self.collection.count()
Expand Down Expand Up @@ -130,15 +137,54 @@ class ChromaDocuments(ChromaCollection):
"""


def _nones_to_none(values):
"""``None`` if every item is ``None`` (chromadb's "field not given")."""
if values is None or all(x is None for x in values):
return None
return values


@appendable(item2kv=uuid_key)
@ValueCodecs.single_nested_value("uris")
class ChromaUris(ChromaCollection):
"""ChromaCollection but reading and writing only the 'uris' field.

Writing uris needs a collection created with a ``data_loader`` (see
``chromadol.data_loaders``); ``chromadb`` refuses uris without one.
``chromadol.data_loaders``): the record is embedded from the data it loads.
``chromadb`` only does that in ``add`` (``upsert`` and ``update`` embed only
documents or images), so a new key is added and an existing key is replaced
(delete, then add, keeping its metadata), and restored if the add fails.
"""

get_include = ("uris",)

def __setitem__(self, k, v: dict):
ids = [k] if isinstance(k, str) else list(k)
old = self.collection.get(
ids, include=["embeddings", "metadatas", "documents", "uris"]
)
if not old["ids"]:
return self.collection.add(ids, **v)
if "metadatas" not in v: # keep metadata, as ``upsert`` would
old_metas = dict(zip(old["ids"], old["metadatas"]))
metas = _nones_to_none([old_metas.get(i) or None for i in ids])
if metas is not None:
v = {**v, "metadatas": metas}
self.collection.delete(old["ids"])
try:
return self.collection.add(ids, **v)
except Exception:
self.collection.add(
ids=old["ids"],
embeddings=old["embeddings"],
# chromadb reads a record without metadata back as ``{}`` but
# refuses ``{}`` on write
metadatas=_nones_to_none([m or None for m in old["metadatas"]]),
documents=_nones_to_none(old["documents"]),
uris=old["uris"],
)
raise


@ValueCodecs.single_nested_value("metadatas")
class ChromaMetadatas(ChromaCollection):
Expand Down
68 changes: 65 additions & 3 deletions chromadol/tests/test_appendable.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,16 +23,28 @@


class _RecordingCollection:
"""Minimal stand-in for a ``chromadb`` Collection, recording ``upsert`` calls."""
"""Minimal stand-in for a ``chromadb`` Collection, recording writes.

``upsert`` and ``add`` are both recorded in ``upserts`` (``ChromaUris`` has
to ``add``, because only ``add`` embeds uris in ``chromadb``).
"""

def __init__(self):
self.upserts = []

def upsert(self, ids, **kwargs):
self.upserts.append((ids, kwargs))

def get(self, ids=None):
return {"ids": [ids for ids, _ in self.upserts]}
add = upsert

def get(self, ids=None, include=None):
wanted = None if ids is None else ([ids] if isinstance(ids, str) else ids)
written = [
i
for recorded, _ in self.upserts
for i in ([recorded] if isinstance(recorded, str) else recorded)
]
return {"ids": [i for i in written if wanted is None or i in wanted]}

def count(self):
return len(self.upserts)
Expand Down Expand Up @@ -133,3 +145,53 @@ def test_chroma_metadatas_reads_the_metadatas_field(tmp_path):
"metadatas": {"author": "me"},
}
assert ChromaMetadatas(collection)["k"] == [{"author": "me"}]


def _uris_store(tmp_path, name):
"""A ``ChromaUris`` over a collection that loads uris as text files."""
from chromadol.data_loaders import FileLoader

client = chromadb.PersistentClient(str(tmp_path / name))
collection = client.create_collection(name, data_loader=FileLoader())
return ChromaUris(collection), collection


def test_chroma_uris_round_trips_against_real_chromadb(tmp_path):
"""Write, read, overwrite and append uris on a real ``chromadb`` collection.

``upsert`` refuses a record with only uris ("Exactly one of documents,
images must be provided"), and a default ``get`` leaves ``uris`` out, so a
store that upserts and reads with the default include can neither write nor
read. A recording fake cannot see either problem.
"""
for name in ("a", "b", "c"):
(tmp_path / f"{name}.txt").write_text(f"contents of {name}")
uris, collection = _uris_store(tmp_path, "uristest")
a, b, c = (str(tmp_path / f"{n}.txt") for n in "abc")

uris["k"] = a
assert uris["k"] == [a]
uris["k"] = c # overwrite a key that has no metadata
assert uris["k"] == [c]

collection.update(ids="k", metadatas={"author": "me"})
uris["k"] = b # overwrite an existing key
assert uris["k"] == [b]
assert collection.get("k")["metadatas"] == [{"author": "me"}]
assert len(uris) == 1

uris.append(c)
assert len(uris) == 2
assert sorted(uris[key][0] for key in uris) == [b, c]


def test_chroma_uris_overwrite_that_fails_keeps_the_old_record(tmp_path):
(tmp_path / "a.txt").write_text("contents of a")
uris, collection = _uris_store(tmp_path, "uriskeep")
a = str(tmp_path / "a.txt")
uris["k"] = a

with pytest.raises(Exception):
uris["k"] = str(tmp_path / "missing.txt") # the data loader can't load it

assert uris["k"] == [a]
Loading