diff --git a/mem0/vector_stores/chroma.py b/mem0/vector_stores/chroma.py index 0d399aad7..12f43cb4b 100644 --- a/mem0/vector_stores/chroma.py +++ b/mem0/vector_stores/chroma.py @@ -185,7 +185,11 @@ class ChromaDB(VectorStoreBase): vector (Optional[List[float]], optional): Updated vector. Defaults to None. payload (Optional[Dict], optional): Updated payload. Defaults to None. """ - self.collection.update(ids=vector_id, embeddings=vector, metadatas=payload) + self.collection.update( + ids=[vector_id], + embeddings=[vector] if vector is not None else None, + metadatas=[payload] if payload is not None else None, + ) def get(self, vector_id: str) -> Optional[OutputData]: """ diff --git a/tests/vector_stores/test_chroma.py b/tests/vector_stores/test_chroma.py index bdea18602..dead40ba4 100644 --- a/tests/vector_stores/test_chroma.py +++ b/tests/vector_stores/test_chroma.py @@ -133,7 +133,31 @@ def test_update_vector(chromadb_instance): chromadb_instance.update(vector_id=vector_id, vector=new_vector, payload=new_payload) chromadb_instance.collection.update.assert_called_once_with( - ids=vector_id, embeddings=new_vector, metadatas=new_payload + ids=[vector_id], embeddings=[new_vector], metadatas=[new_payload] + ) + + +def test_update_vector_metadata_only(chromadb_instance): + # Metadata-only update (vector=None) must not wrap None in a list. + vector_id = "id1" + new_payload = {"name": "updated_vector"} + + chromadb_instance.update(vector_id=vector_id, vector=None, payload=new_payload) + + chromadb_instance.collection.update.assert_called_once_with( + ids=[vector_id], embeddings=None, metadatas=[new_payload] + ) + + +def test_update_vector_embedding_only(chromadb_instance): + # Vector-only update (payload=None) must not wrap None in a list. + vector_id = "id1" + new_vector = [0.7, 0.8, 0.9] + + chromadb_instance.update(vector_id=vector_id, vector=new_vector, payload=None) + + chromadb_instance.collection.update.assert_called_once_with( + ids=[vector_id], embeddings=[new_vector], metadatas=None )