diff --git a/chromadol/base.py b/chromadol/base.py index cac11bb..6f18074 100644 --- a/chromadol/base.py +++ b/chromadol/base.py @@ -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() @@ -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): diff --git a/chromadol/tests/test_appendable.py b/chromadol/tests/test_appendable.py index 891b439..795d6be 100644 --- a/chromadol/tests/test_appendable.py +++ b/chromadol/tests/test_appendable.py @@ -23,7 +23,11 @@ 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 = [] @@ -31,8 +35,16 @@ def __init__(self): 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) @@ -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]