Added mulit id filters support for all vectorstores (#3269)

This commit is contained in:
Parshva Daftari
2025-08-05 03:22:57 +05:30
committed by GitHub
parent 57a16aeb4b
commit 42e60d6724
8 changed files with 728 additions and 45 deletions
+27 -2
View File
@@ -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}
+3 -2
View File
@@ -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)
+20 -1
View File
@@ -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}'.")
+4 -1
View File
@@ -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."""
+145
View File
@@ -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
@@ -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 == []
+203 -38
View File
@@ -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
+183 -1
View File
@@ -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")