From 42e60d672448ee69c500d466a7e4dc0cf607f22c Mon Sep 17 00:00:00 2001 From: Parshva Daftari <89991302+parshvadaftari@users.noreply.github.com> Date: Tue, 5 Aug 2025 03:22:57 +0530 Subject: [PATCH] Added mulit id filters support for all vectorstores (#3269) --- mem0/vector_stores/chroma.py | 29 ++- mem0/vector_stores/langchain.py | 5 +- mem0/vector_stores/mongodb.py | 21 +- mem0/vector_stores/qdrant.py | 5 +- tests/vector_stores/test_chroma.py | 145 +++++++++++ .../test_langchain_vector_store.py | 143 +++++++++++ tests/vector_stores/test_mongodb.py | 241 +++++++++++++++--- tests/vector_stores/test_qdrant.py | 184 ++++++++++++- 8 files changed, 728 insertions(+), 45 deletions(-) diff --git a/mem0/vector_stores/chroma.py b/mem0/vector_stores/chroma.py index 1de95ad1a..681d4626c 100644 --- a/mem0/vector_stores/chroma.py +++ b/mem0/vector_stores/chroma.py @@ -142,7 +142,8 @@ class ChromaDB(VectorStoreBase): Returns: List[OutputData]: Search results. """ - results = self.collection.query(query_embeddings=vectors, where=filters, n_results=limit) + where_clause = self._generate_where_clause(filters) if filters else None + results = self.collection.query(query_embeddings=vectors, where=where_clause, n_results=limit) final_results = self._parse_output(results) return final_results @@ -219,7 +220,8 @@ class ChromaDB(VectorStoreBase): Returns: List[OutputData]: List of vectors. """ - results = self.collection.get(where=filters, limit=limit) + where_clause = self._generate_where_clause(filters) if filters else None + results = self.collection.get(where=where_clause, limit=limit) return [self._parse_output(results)] def reset(self): @@ -227,3 +229,26 @@ class ChromaDB(VectorStoreBase): logger.warning(f"Resetting index {self.collection_name}...") self.delete_col() self.collection = self.create_col(self.collection_name) + + @staticmethod + def _generate_where_clause(where: dict[str, any]) -> dict[str, any]: + """ + Generate a properly formatted where clause for ChromaDB. + + Args: + where (dict[str, any]): The filter conditions. + + Returns: + dict[str, any]: Properly formatted where clause for ChromaDB. + """ + # If only one filter is supplied, return it as is + # (no need to wrap in $and based on chroma docs) + if where is None: + return {} + if len(where.keys()) <= 1: + return where + where_filters = [] + for k, v in where.items(): + if isinstance(v, str): + where_filters.append({k: v}) + return {"$and": where_filters} diff --git a/mem0/vector_stores/langchain.py b/mem0/vector_stores/langchain.py index f3fcf07e8..4fe06c1b1 100644 --- a/mem0/vector_stores/langchain.py +++ b/mem0/vector_stores/langchain.py @@ -160,8 +160,9 @@ class Langchain(VectorStoreBase): if hasattr(self.client, "_collection") and hasattr(self.client._collection, "get"): # Convert mem0 filters to Chroma where clause if needed where_clause = None - if filters and "user_id" in filters: - where_clause = {"user_id": filters["user_id"]} + if filters: + # Handle all filters, not just user_id + where_clause = filters result = self.client._collection.get(where=where_clause, limit=limit) diff --git a/mem0/vector_stores/mongodb.py b/mem0/vector_stores/mongodb.py index ef82c6e60..381bfde4e 100644 --- a/mem0/vector_stores/mongodb.py +++ b/mem0/vector_stores/mongodb.py @@ -148,6 +148,17 @@ class MongoDB(VectorStoreBase): {"$set": {"score": {"$meta": "vectorSearchScore"}}}, {"$project": {"embedding": 0}}, ] + + # Add filter stage if filters are provided + if filters: + filter_conditions = [] + for key, value in filters.items(): + filter_conditions.append({"payload." + key: value}) + + if filter_conditions: + # Add a $match stage after vector search to apply filters + pipeline.insert(1, {"$match": {"$and": filter_conditions}}) + results = list(collection.aggregate(pipeline)) logger.info(f"Vector search completed. Found {len(results)} documents.") except Exception as e: @@ -271,7 +282,15 @@ class MongoDB(VectorStoreBase): List[OutputData]: List of vectors. """ try: - query = filters or {} + query = {} + if filters: + # Apply filters to the payload field + filter_conditions = [] + for key, value in filters.items(): + filter_conditions.append({"payload." + key: value}) + if filter_conditions: + query = {"$and": filter_conditions} + cursor = self.collection.find(query).limit(limit) results = [OutputData(id=str(doc["_id"]), score=None, payload=doc.get("payload")) for doc in cursor] logger.info(f"Retrieved {len(results)} documents from collection '{self.collection_name}'.") diff --git a/mem0/vector_stores/qdrant.py b/mem0/vector_stores/qdrant.py index 98604398c..273b61428 100644 --- a/mem0/vector_stores/qdrant.py +++ b/mem0/vector_stores/qdrant.py @@ -148,6 +148,9 @@ class Qdrant(VectorStoreBase): Returns: Filter: The created Filter object. """ + if not filters: + return None + conditions = [] for key, value in filters.items(): if isinstance(value, dict) and "gte" in value and "lte" in value: @@ -258,7 +261,7 @@ class Qdrant(VectorStoreBase): with_payload=True, with_vectors=False, ) - return result + return result.points def reset(self): """Reset the index by deleting and recreating it.""" diff --git a/tests/vector_stores/test_chroma.py b/tests/vector_stores/test_chroma.py index 39ddfdca7..4339c8c74 100644 --- a/tests/vector_stores/test_chroma.py +++ b/tests/vector_stores/test_chroma.py @@ -48,6 +48,73 @@ def test_search_vectors(chromadb_instance, mock_chromadb_client): assert results[0].payload == {"name": "vector1"} +def test_search_vectors_with_filters(chromadb_instance, mock_chromadb_client): + """Test search with agent_id and run_id filters.""" + mock_result = { + "ids": [["id1"]], + "distances": [[0.1]], + "metadatas": [[{"name": "vector1", "user_id": "alice", "agent_id": "agent1", "run_id": "run1"}]], + } + chromadb_instance.collection.query.return_value = mock_result + + vectors = [[0.1, 0.2, 0.3]] + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = chromadb_instance.search(query="", vectors=vectors, limit=2, filters=filters) + + # Verify that _generate_where_clause was called with the filters + expected_where = {"$and": [{"user_id": "alice"}, {"agent_id": "agent1"}, {"run_id": "run1"}]} + chromadb_instance.collection.query.assert_called_once_with( + query_embeddings=vectors, where=expected_where, n_results=2 + ) + + assert len(results) == 1 + assert results[0].id == "id1" + assert results[0].payload["user_id"] == "alice" + assert results[0].payload["agent_id"] == "agent1" + assert results[0].payload["run_id"] == "run1" + + +def test_search_vectors_with_single_filter(chromadb_instance, mock_chromadb_client): + """Test search with single filter (should not use $and).""" + mock_result = { + "ids": [["id1"]], + "distances": [[0.1]], + "metadatas": [[{"name": "vector1", "user_id": "alice"}]], + } + chromadb_instance.collection.query.return_value = mock_result + + vectors = [[0.1, 0.2, 0.3]] + filters = {"user_id": "alice"} + results = chromadb_instance.search(query="", vectors=vectors, limit=2, filters=filters) + + # Verify that single filter is passed as-is (no $and wrapper) + chromadb_instance.collection.query.assert_called_once_with( + query_embeddings=vectors, where=filters, n_results=2 + ) + + assert len(results) == 1 + assert results[0].payload["user_id"] == "alice" + + +def test_search_vectors_with_no_filters(chromadb_instance, mock_chromadb_client): + """Test search with no filters.""" + mock_result = { + "ids": [["id1"]], + "distances": [[0.1]], + "metadatas": [[{"name": "vector1"}]], + } + chromadb_instance.collection.query.return_value = mock_result + + vectors = [[0.1, 0.2, 0.3]] + results = chromadb_instance.search(query="", vectors=vectors, limit=2, filters=None) + + chromadb_instance.collection.query.assert_called_once_with( + query_embeddings=vectors, where=None, n_results=2 + ) + + assert len(results) == 1 + + def test_delete_vector(chromadb_instance): vector_id = "id1" @@ -100,3 +167,81 @@ def test_list_vectors(chromadb_instance): assert len(results[0]) == 2 assert results[0][0].id == "id1" assert results[0][1].id == "id2" + + +def test_list_vectors_with_filters(chromadb_instance): + """Test list with agent_id and run_id filters.""" + mock_result = { + "ids": [["id1"]], + "distances": [[0.1]], + "metadatas": [[{"name": "vector1", "user_id": "alice", "agent_id": "agent1", "run_id": "run1"}]], + } + chromadb_instance.collection.get.return_value = mock_result + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = chromadb_instance.list(filters=filters, limit=2) + + # Verify that _generate_where_clause was called with the filters + expected_where = {"$and": [{"user_id": "alice"}, {"agent_id": "agent1"}, {"run_id": "run1"}]} + chromadb_instance.collection.get.assert_called_once_with(where=expected_where, limit=2) + + assert len(results[0]) == 1 + assert results[0][0].payload["user_id"] == "alice" + assert results[0][0].payload["agent_id"] == "agent1" + assert results[0][0].payload["run_id"] == "run1" + + +def test_list_vectors_with_single_filter(chromadb_instance): + """Test list with single filter (should not use $and).""" + mock_result = { + "ids": [["id1"]], + "distances": [[0.1]], + "metadatas": [[{"name": "vector1", "user_id": "alice"}]], + } + chromadb_instance.collection.get.return_value = mock_result + + filters = {"user_id": "alice"} + results = chromadb_instance.list(filters=filters, limit=2) + + # Verify that single filter is passed as-is (no $and wrapper) + chromadb_instance.collection.get.assert_called_once_with(where=filters, limit=2) + + assert len(results[0]) == 1 + assert results[0][0].payload["user_id"] == "alice" + + +def test_generate_where_clause_multiple_filters(): + """Test _generate_where_clause with multiple filters.""" + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + result = ChromaDB._generate_where_clause(filters) + + expected = {"$and": [{"user_id": "alice"}, {"agent_id": "agent1"}, {"run_id": "run1"}]} + assert result == expected + + +def test_generate_where_clause_single_filter(): + """Test _generate_where_clause with single filter.""" + filters = {"user_id": "alice"} + result = ChromaDB._generate_where_clause(filters) + + # Single filter should be returned as-is + assert result == filters + + +def test_generate_where_clause_no_filters(): + """Test _generate_where_clause with no filters.""" + result = ChromaDB._generate_where_clause(None) + assert result == {} + + result = ChromaDB._generate_where_clause({}) + assert result == {} + + +def test_generate_where_clause_non_string_values(): + """Test _generate_where_clause with non-string values.""" + filters = {"user_id": "alice", "count": 5, "active": True} + result = ChromaDB._generate_where_clause(filters) + + # Only string values should be included in $and array + expected = {"$and": [{"user_id": "alice"}]} + assert result == expected diff --git a/tests/vector_stores/test_langchain_vector_store.py b/tests/vector_stores/test_langchain_vector_store.py index 6e156ec22..5f1a8c544 100644 --- a/tests/vector_stores/test_langchain_vector_store.py +++ b/tests/vector_stores/test_langchain_vector_store.py @@ -64,6 +64,66 @@ def test_search_vectors(langchain_instance): langchain_instance.client.similarity_search_by_vector.assert_called_with(embedding=vectors, k=2, filter=filters) +def test_search_vectors_with_agent_id_run_id_filters(langchain_instance): + """Test search with agent_id and run_id filters.""" + # Mock search results + mock_docs = [ + Mock(metadata={"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}, id="id1"), + Mock(metadata={"user_id": "bob", "agent_id": "agent2", "run_id": "run2"}, id="id2") + ] + langchain_instance.client.similarity_search_by_vector.return_value = mock_docs + + vectors = [[0.1, 0.2, 0.3]] + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = langchain_instance.search(query="", vectors=vectors, limit=2, filters=filters) + + # Verify that filters were passed to the underlying vector store + langchain_instance.client.similarity_search_by_vector.assert_called_once_with( + embedding=vectors, k=2, filter=filters + ) + + assert len(results) == 2 + assert results[0].payload["user_id"] == "alice" + assert results[0].payload["agent_id"] == "agent1" + assert results[0].payload["run_id"] == "run1" + + +def test_search_vectors_with_single_filter(langchain_instance): + """Test search with single filter.""" + # Mock search results + mock_docs = [Mock(metadata={"user_id": "alice"}, id="id1")] + langchain_instance.client.similarity_search_by_vector.return_value = mock_docs + + vectors = [[0.1, 0.2, 0.3]] + filters = {"user_id": "alice"} + results = langchain_instance.search(query="", vectors=vectors, limit=2, filters=filters) + + # Verify that filters were passed to the underlying vector store + langchain_instance.client.similarity_search_by_vector.assert_called_once_with( + embedding=vectors, k=2, filter=filters + ) + + assert len(results) == 1 + assert results[0].payload["user_id"] == "alice" + + +def test_search_vectors_with_no_filters(langchain_instance): + """Test search with no filters.""" + # Mock search results + mock_docs = [Mock(metadata={"name": "vector1"}, id="id1")] + langchain_instance.client.similarity_search_by_vector.return_value = mock_docs + + vectors = [[0.1, 0.2, 0.3]] + results = langchain_instance.search(query="", vectors=vectors, limit=2, filters=None) + + # Verify that no filters were passed to the underlying vector store + langchain_instance.client.similarity_search_by_vector.assert_called_once_with( + embedding=vectors, k=2 + ) + + assert len(results) == 1 + + def test_get_vector(langchain_instance): # Mock get result mock_doc = Mock(metadata={"name": "vector1"}, id="id1") @@ -81,3 +141,86 @@ def test_get_vector(langchain_instance): langchain_instance.client.get_by_ids.return_value = [] result = langchain_instance.get("non_existent_id") assert result is None + + +def test_list_with_filters(langchain_instance): + """Test list with agent_id and run_id filters.""" + # Mock the _collection.get method + mock_collection = Mock() + mock_collection.get.return_value = { + "ids": [["id1"]], + "metadatas": [[{"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}]], + "documents": [["test document"]] + } + langchain_instance.client._collection = mock_collection + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = langchain_instance.list(filters=filters, limit=10) + + # Verify that the collection.get method was called with the correct filters + mock_collection.get.assert_called_once_with(where=filters, limit=10) + + # Verify the results + assert len(results) == 1 + assert len(results[0]) == 1 + assert results[0][0].payload["user_id"] == "alice" + assert results[0][0].payload["agent_id"] == "agent1" + assert results[0][0].payload["run_id"] == "run1" + + +def test_list_with_single_filter(langchain_instance): + """Test list with single filter.""" + # Mock the _collection.get method + mock_collection = Mock() + mock_collection.get.return_value = { + "ids": [["id1"]], + "metadatas": [[{"user_id": "alice"}]], + "documents": [["test document"]] + } + langchain_instance.client._collection = mock_collection + + filters = {"user_id": "alice"} + results = langchain_instance.list(filters=filters, limit=10) + + # Verify that the collection.get method was called with the correct filter + mock_collection.get.assert_called_once_with(where=filters, limit=10) + + # Verify the results + assert len(results) == 1 + assert len(results[0]) == 1 + assert results[0][0].payload["user_id"] == "alice" + + +def test_list_with_no_filters(langchain_instance): + """Test list with no filters.""" + # Mock the _collection.get method + mock_collection = Mock() + mock_collection.get.return_value = { + "ids": [["id1"]], + "metadatas": [[{"name": "vector1"}]], + "documents": [["test document"]] + } + langchain_instance.client._collection = mock_collection + + results = langchain_instance.list(filters=None, limit=10) + + # Verify that the collection.get method was called with no filters + mock_collection.get.assert_called_once_with(where=None, limit=10) + + # Verify the results + assert len(results) == 1 + assert len(results[0]) == 1 + assert results[0][0].payload["name"] == "vector1" + + +def test_list_with_exception(langchain_instance): + """Test list when an exception occurs.""" + # Mock the _collection.get method to raise an exception + mock_collection = Mock() + mock_collection.get.side_effect = Exception("Test exception") + langchain_instance.client._collection = mock_collection + + results = langchain_instance.list(filters={"user_id": "alice"}, limit=10) + + # Verify that an empty list is returned when an exception occurs + assert results == [] diff --git a/tests/vector_stores/test_mongodb.py b/tests/vector_stores/test_mongodb.py index abb299e01..bb4b3f69f 100644 --- a/tests/vector_stores/test_mongodb.py +++ b/tests/vector_stores/test_mongodb.py @@ -14,7 +14,12 @@ def mongo_vector_fixture(mock_mongo_client): mock_collection.list_search_indexes.return_value = [] mock_collection.aggregate.return_value = [] mock_collection.find_one.return_value = None - mock_collection.find.return_value = [] + + # Create a proper mock cursor + mock_cursor = MagicMock() + mock_cursor.limit.return_value = mock_cursor + mock_collection.find.return_value = mock_cursor + mock_db.list_collection_names.return_value = [] mongo_vector = MongoDB( @@ -102,87 +107,247 @@ def test_search(mongo_vector_fixture): assert len(results) == 2 assert results[0].id == "id1" assert results[0].score == 0.9 - assert results[1].id == "id2" - assert results[1].score == 0.8 + assert results[0].payload == {"key": "value1"} + + +def test_search_with_filters(mongo_vector_fixture): + """Test search with agent_id and run_id filters.""" + mongo_vector, mock_collection, _ = mongo_vector_fixture + query_vector = [0.1] * 1536 + mock_collection.aggregate.return_value = [ + {"_id": "id1", "score": 0.9, "payload": {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}}, + ] + mock_collection.list_search_indexes.return_value = ["test_collection_vector_index"] + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = mongo_vector.search("query_str", query_vector, limit=2, filters=filters) + + # Verify that the aggregation pipeline includes the filter stage + mock_collection.aggregate.assert_called_once() + pipeline = mock_collection.aggregate.call_args[0][0] + + # Check that the pipeline has the expected stages + assert len(pipeline) == 4 # vectorSearch, match, set, project + + # Check that the match stage is present with the correct filters + match_stage = pipeline[1] + assert "$match" in match_stage + assert match_stage["$match"]["$and"] == [ + {"payload.user_id": "alice"}, + {"payload.agent_id": "agent1"}, + {"payload.run_id": "run1"} + ] + + assert len(results) == 1 + assert results[0].payload["user_id"] == "alice" + assert results[0].payload["agent_id"] == "agent1" + assert results[0].payload["run_id"] == "run1" + + +def test_search_with_single_filter(mongo_vector_fixture): + """Test search with single filter.""" + mongo_vector, mock_collection, _ = mongo_vector_fixture + query_vector = [0.1] * 1536 + mock_collection.aggregate.return_value = [ + {"_id": "id1", "score": 0.9, "payload": {"user_id": "alice"}}, + ] + mock_collection.list_search_indexes.return_value = ["test_collection_vector_index"] + + filters = {"user_id": "alice"} + results = mongo_vector.search("query_str", query_vector, limit=2, filters=filters) + + # Verify that the aggregation pipeline includes the filter stage + mock_collection.aggregate.assert_called_once() + pipeline = mock_collection.aggregate.call_args[0][0] + + # Check that the match stage is present with the correct filter + match_stage = pipeline[1] + assert "$match" in match_stage + assert match_stage["$match"]["$and"] == [{"payload.user_id": "alice"}] + + assert len(results) == 1 + assert results[0].payload["user_id"] == "alice" + + +def test_search_with_no_filters(mongo_vector_fixture): + """Test search with no filters.""" + mongo_vector, mock_collection, _ = mongo_vector_fixture + query_vector = [0.1] * 1536 + mock_collection.aggregate.return_value = [ + {"_id": "id1", "score": 0.9, "payload": {"key": "value1"}}, + ] + mock_collection.list_search_indexes.return_value = ["test_collection_vector_index"] + + results = mongo_vector.search("query_str", query_vector, limit=2, filters=None) + + # Verify that the aggregation pipeline does not include the filter stage + mock_collection.aggregate.assert_called_once() + pipeline = mock_collection.aggregate.call_args[0][0] + + # Check that the pipeline has only the expected stages (no match stage) + assert len(pipeline) == 3 # vectorSearch, set, project + + assert len(results) == 1 def test_delete(mongo_vector_fixture): mongo_vector, mock_collection, _ = mongo_vector_fixture - mock_delete_result = MagicMock() - mock_delete_result.deleted_count = 1 - mock_collection.delete_one.return_value = mock_delete_result + vector_id = "id1" + mock_collection.delete_one.return_value = MagicMock(deleted_count=1) + + # Reset the mock to clear calls from fixture setup + mock_collection.delete_one.reset_mock() - mongo_vector.delete("id1") - mock_collection.delete_one.assert_called_with({"_id": "id1"}) + mongo_vector.delete(vector_id=vector_id) + + mock_collection.delete_one.assert_called_once_with({"_id": vector_id}) def test_update(mongo_vector_fixture): mongo_vector, mock_collection, _ = mongo_vector_fixture - mock_update_result = MagicMock() - mock_update_result.matched_count = 1 - mock_collection.update_one.return_value = mock_update_result - idValue = "id1" - vectorValue = [0.2] * 1536 - payloadValue = {"key": "updated"} + vector_id = "id1" + updated_vector = [0.3] * 1536 + updated_payload = {"name": "updated_vector"} + + mock_collection.update_one.return_value = MagicMock(matched_count=1) + + mongo_vector.update(vector_id=vector_id, vector=updated_vector, payload=updated_payload) - mongo_vector.update(idValue, vector=vectorValue, payload=payloadValue) mock_collection.update_one.assert_called_once_with( - {"_id": idValue}, - {"$set": {"embedding": vectorValue, "payload": payloadValue}}, + {"_id": vector_id}, {"$set": {"embedding": updated_vector, "payload": updated_payload}} ) def test_get(mongo_vector_fixture): mongo_vector, mock_collection, _ = mongo_vector_fixture - mock_collection.find_one.return_value = {"_id": "id1", "payload": {"key": "value1"}} + vector_id = "id1" + mock_collection.find_one.return_value = {"_id": vector_id, "payload": {"key": "value"}} - result = mongo_vector.get("id1") - assert result is not None - assert result.id == "id1" - assert result.payload == {"key": "value1"} + result = mongo_vector.get(vector_id=vector_id) + + mock_collection.find_one.assert_called_once_with({"_id": vector_id}) + assert result.id == vector_id + assert result.payload == {"key": "value"} def test_list_cols(mongo_vector_fixture): mongo_vector, _, mock_db = mongo_vector_fixture - mock_db.list_collection_names.return_value = ["col1", "col2"] + mock_db.list_collection_names.return_value = ["collection1", "collection2"] + + # Reset the mock to clear calls from fixture setup + mock_db.list_collection_names.reset_mock() - collections = mongo_vector.list_cols() - assert collections == ["col1", "col2"] + result = mongo_vector.list_cols() + + mock_db.list_collection_names.assert_called_once() + assert result == ["collection1", "collection2"] def test_delete_col(mongo_vector_fixture): mongo_vector, mock_collection, _ = mongo_vector_fixture mongo_vector.delete_col() + mock_collection.drop.assert_called_once() def test_col_info(mongo_vector_fixture): - mongo_vector, _, mock_db = mongo_vector_fixture + mongo_vector, mock_collection, mock_db = mongo_vector_fixture mock_db.command.return_value = {"count": 10, "size": 1024} - info = mongo_vector.col_info() + result = mongo_vector.col_info() + mock_db.command.assert_called_once_with("collstats", "test_collection") - assert info["name"] == "test_collection" - assert info["count"] == 10 - assert info["size"] == 1024 + assert result["name"] == "test_collection" + assert result["count"] == 10 + assert result["size"] == 1024 def test_list(mongo_vector_fixture): mongo_vector, mock_collection, _ = mongo_vector_fixture - mock_cursor = MagicMock() - mock_cursor.limit.return_value = [ + # Mock the cursor to return the expected data + mock_cursor = mock_collection.find.return_value + mock_cursor.__iter__.return_value = [ {"_id": "id1", "payload": {"key": "value1"}}, {"_id": "id2", "payload": {"key": "value2"}}, ] - mock_collection.find.return_value = mock_cursor - query_filters = {"_id": {"$in": ["id1", "id2"]}} - results = mongo_vector.list(filters=query_filters, limit=2) - mock_collection.find.assert_called_once_with(query_filters) + results = mongo_vector.list(limit=2) + + mock_collection.find.assert_called_once_with({}) mock_cursor.limit.assert_called_once_with(2) assert len(results) == 2 assert results[0].id == "id1" assert results[0].payload == {"key": "value1"} - assert results[1].id == "id2" - assert results[1].payload == {"key": "value2"} + + +def test_list_with_filters(mongo_vector_fixture): + """Test list with agent_id and run_id filters.""" + mongo_vector, mock_collection, _ = mongo_vector_fixture + # Mock the cursor to return the expected data + mock_cursor = mock_collection.find.return_value + mock_cursor.__iter__.return_value = [ + {"_id": "id1", "payload": {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}}, + ] + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = mongo_vector.list(filters=filters, limit=2) + + # Verify that the find method was called with the correct query + expected_query = { + "$and": [ + {"payload.user_id": "alice"}, + {"payload.agent_id": "agent1"}, + {"payload.run_id": "run1"} + ] + } + mock_collection.find.assert_called_once_with(expected_query) + mock_cursor.limit.assert_called_once_with(2) + + assert len(results) == 1 + assert results[0].payload["user_id"] == "alice" + assert results[0].payload["agent_id"] == "agent1" + assert results[0].payload["run_id"] == "run1" + + +def test_list_with_single_filter(mongo_vector_fixture): + """Test list with single filter.""" + mongo_vector, mock_collection, _ = mongo_vector_fixture + # Mock the cursor to return the expected data + mock_cursor = mock_collection.find.return_value + mock_cursor.__iter__.return_value = [ + {"_id": "id1", "payload": {"user_id": "alice"}}, + ] + + filters = {"user_id": "alice"} + results = mongo_vector.list(filters=filters, limit=2) + + # Verify that the find method was called with the correct query + expected_query = { + "$and": [ + {"payload.user_id": "alice"} + ] + } + mock_collection.find.assert_called_once_with(expected_query) + mock_cursor.limit.assert_called_once_with(2) + + assert len(results) == 1 + assert results[0].payload["user_id"] == "alice" + + +def test_list_with_no_filters(mongo_vector_fixture): + """Test list with no filters.""" + mongo_vector, mock_collection, _ = mongo_vector_fixture + # Mock the cursor to return the expected data + mock_cursor = mock_collection.find.return_value + mock_cursor.__iter__.return_value = [ + {"_id": "id1", "payload": {"key": "value1"}}, + ] + + results = mongo_vector.list(filters=None, limit=2) + + # Verify that the find method was called with empty query + mock_collection.find.assert_called_once_with({}) + mock_cursor.limit.assert_called_once_with(2) + + assert len(results) == 1 diff --git a/tests/vector_stores/test_qdrant.py b/tests/vector_stores/test_qdrant.py index 0a2bb444a..dfcf145dc 100644 --- a/tests/vector_stores/test_qdrant.py +++ b/tests/vector_stores/test_qdrant.py @@ -3,7 +3,13 @@ import uuid from unittest.mock import MagicMock from qdrant_client import QdrantClient -from qdrant_client.models import Distance, PointIdsList, PointStruct, VectorParams +from qdrant_client.models import ( + Distance, + Filter, + PointIdsList, + PointStruct, + VectorParams, +) from mem0.vector_stores.qdrant import Qdrant @@ -64,6 +70,121 @@ class TestQdrant(unittest.TestCase): self.assertEqual(results[0].payload, {"key": "value"}) self.assertEqual(results[0].score, 0.95) + def test_search_with_filters(self): + """Test search with agent_id and run_id filters.""" + vectors = [[0.1, 0.2]] + mock_point = MagicMock( + id=str(uuid.uuid4()), + score=0.95, + payload={"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + ) + self.client_mock.query_points.return_value = MagicMock(points=[mock_point]) + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = self.qdrant.search(query="", vectors=vectors, limit=1, filters=filters) + + # Verify that _create_filter was called and query_filter was passed + self.client_mock.query_points.assert_called_once() + call_args = self.client_mock.query_points.call_args[1] + self.assertEqual(call_args["collection_name"], "test_collection") + self.assertEqual(call_args["query"], vectors) + self.assertEqual(call_args["limit"], 1) + + # Verify that a Filter object was created + query_filter = call_args["query_filter"] + self.assertIsInstance(query_filter, Filter) + self.assertEqual(len(query_filter.must), 3) # user_id, agent_id, run_id + + self.assertEqual(len(results), 1) + self.assertEqual(results[0].payload["user_id"], "alice") + self.assertEqual(results[0].payload["agent_id"], "agent1") + self.assertEqual(results[0].payload["run_id"], "run1") + + def test_search_with_single_filter(self): + """Test search with single filter.""" + vectors = [[0.1, 0.2]] + mock_point = MagicMock( + id=str(uuid.uuid4()), + score=0.95, + payload={"user_id": "alice"} + ) + self.client_mock.query_points.return_value = MagicMock(points=[mock_point]) + + filters = {"user_id": "alice"} + results = self.qdrant.search(query="", vectors=vectors, limit=1, filters=filters) + + # Verify that a Filter object was created with single condition + call_args = self.client_mock.query_points.call_args[1] + query_filter = call_args["query_filter"] + self.assertIsInstance(query_filter, Filter) + self.assertEqual(len(query_filter.must), 1) # Only user_id + + self.assertEqual(len(results), 1) + self.assertEqual(results[0].payload["user_id"], "alice") + + def test_search_with_no_filters(self): + """Test search with no filters.""" + vectors = [[0.1, 0.2]] + mock_point = MagicMock(id=str(uuid.uuid4()), score=0.95, payload={"key": "value"}) + self.client_mock.query_points.return_value = MagicMock(points=[mock_point]) + + results = self.qdrant.search(query="", vectors=vectors, limit=1, filters=None) + + call_args = self.client_mock.query_points.call_args[1] + self.assertIsNone(call_args["query_filter"]) + + self.assertEqual(len(results), 1) + + def test_create_filter_multiple_filters(self): + """Test _create_filter with multiple filters.""" + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + result = self.qdrant._create_filter(filters) + + self.assertIsInstance(result, Filter) + self.assertEqual(len(result.must), 3) + + # Check that all conditions are present + conditions = [cond.key for cond in result.must] + self.assertIn("user_id", conditions) + self.assertIn("agent_id", conditions) + self.assertIn("run_id", conditions) + + def test_create_filter_single_filter(self): + """Test _create_filter with single filter.""" + filters = {"user_id": "alice"} + result = self.qdrant._create_filter(filters) + + self.assertIsInstance(result, Filter) + self.assertEqual(len(result.must), 1) + self.assertEqual(result.must[0].key, "user_id") + self.assertEqual(result.must[0].match.value, "alice") + + def test_create_filter_no_filters(self): + """Test _create_filter with no filters.""" + result = self.qdrant._create_filter(None) + self.assertIsNone(result) + + result = self.qdrant._create_filter({}) + self.assertIsNone(result) + + def test_create_filter_with_range_values(self): + """Test _create_filter with range values.""" + filters = {"user_id": "alice", "count": {"gte": 5, "lte": 10}} + result = self.qdrant._create_filter(filters) + + self.assertIsInstance(result, Filter) + self.assertEqual(len(result.must), 2) + + # Check that range condition is created + range_conditions = [cond for cond in result.must if hasattr(cond, 'range') and cond.range is not None] + self.assertEqual(len(range_conditions), 1) + self.assertEqual(range_conditions[0].key, "count") + + # Check that string condition is created + string_conditions = [cond for cond in result.must if hasattr(cond, 'match') and cond.match is not None] + self.assertEqual(len(string_conditions), 1) + self.assertEqual(string_conditions[0].key, "user_id") + def test_delete(self): vector_id = str(uuid.uuid4()) self.qdrant.delete(vector_id=vector_id) @@ -103,6 +224,67 @@ class TestQdrant(unittest.TestCase): result = self.qdrant.list_cols() self.assertEqual(result.collections[0]["name"], "test_collection") + def test_list_with_filters(self): + """Test list with agent_id and run_id filters.""" + mock_point = MagicMock( + id=str(uuid.uuid4()), + score=0.95, + payload={"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + ) + self.client_mock.scroll.return_value = MagicMock(points=[mock_point]) + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = self.qdrant.list(filters=filters, limit=10) + + # Verify that _create_filter was called and scroll_filter was passed + self.client_mock.scroll.assert_called_once() + call_args = self.client_mock.scroll.call_args[1] + self.assertEqual(call_args["collection_name"], "test_collection") + self.assertEqual(call_args["limit"], 10) + + # Verify that a Filter object was created + scroll_filter = call_args["scroll_filter"] + self.assertIsInstance(scroll_filter, Filter) + self.assertEqual(len(scroll_filter.must), 3) # user_id, agent_id, run_id + + self.assertEqual(len(results), 1) + self.assertEqual(results[0].payload["user_id"], "alice") + self.assertEqual(results[0].payload["agent_id"], "agent1") + self.assertEqual(results[0].payload["run_id"], "run1") + + def test_list_with_single_filter(self): + """Test list with single filter.""" + mock_point = MagicMock( + id=str(uuid.uuid4()), + score=0.95, + payload={"user_id": "alice"} + ) + self.client_mock.scroll.return_value = MagicMock(points=[mock_point]) + + filters = {"user_id": "alice"} + results = self.qdrant.list(filters=filters, limit=10) + + # Verify that a Filter object was created with single condition + call_args = self.client_mock.scroll.call_args[1] + scroll_filter = call_args["scroll_filter"] + self.assertIsInstance(scroll_filter, Filter) + self.assertEqual(len(scroll_filter.must), 1) # Only user_id + + self.assertEqual(len(results), 1) + self.assertEqual(results[0].payload["user_id"], "alice") + + def test_list_with_no_filters(self): + """Test list with no filters.""" + mock_point = MagicMock(id=str(uuid.uuid4()), score=0.95, payload={"key": "value"}) + self.client_mock.scroll.return_value = MagicMock(points=[mock_point]) + + results = self.qdrant.list(filters=None, limit=10) + + call_args = self.client_mock.scroll.call_args[1] + self.assertIsNone(call_args["scroll_filter"]) + + self.assertEqual(len(results), 1) + def test_delete_col(self): self.qdrant.delete_col() self.client_mock.delete_collection.assert_called_once_with(collection_name="test_collection")