fix(vector-stores): filter S3 vector list results (#5018)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -182,10 +182,6 @@ class S3Vectors(VectorStoreBase):
|
||||
return response.get("index", {})
|
||||
|
||||
def list(self, filters=None, top_k=None):
|
||||
# Note: list_vectors does not support metadata filtering.
|
||||
if filters:
|
||||
logger.warning("S3 Vectors `list` does not support metadata filtering. Ignoring filters.")
|
||||
|
||||
params = {
|
||||
"vectorBucketName": self.vector_bucket_name,
|
||||
"indexName": self.collection_name,
|
||||
@@ -200,7 +196,16 @@ class S3Vectors(VectorStoreBase):
|
||||
all_vectors = []
|
||||
for page in pages:
|
||||
all_vectors.extend(page.get("vectors", []))
|
||||
return [self._parse_output(all_vectors)]
|
||||
results = self._parse_output(all_vectors)
|
||||
if filters:
|
||||
results = [
|
||||
result
|
||||
for result in results
|
||||
if result.payload and all(result.payload.get(k) == v for k, v in filters.items())
|
||||
]
|
||||
if top_k:
|
||||
results = results[:top_k]
|
||||
return [results]
|
||||
|
||||
def reset(self):
|
||||
logger.warning(f"Resetting index {self.collection_name}...")
|
||||
|
||||
@@ -223,3 +223,26 @@ def test_reset(mock_boto_client):
|
||||
vectorBucketName=BUCKET_NAME, indexName=INDEX_NAME
|
||||
)
|
||||
assert mock_boto_client.create_index.call_count == 2
|
||||
|
||||
|
||||
def test_list_filters_metadata_client_side(mock_boto_client):
|
||||
"""S3 list_vectors does not support metadata filters, so the store filters returned payloads locally."""
|
||||
mock_paginator = mock_boto_client.get_paginator.return_value
|
||||
mock_paginator.paginate.return_value = [
|
||||
{
|
||||
"vectors": [
|
||||
{"key": "id1", "metadata": {"user_id": "alice", "category": "work"}},
|
||||
{"key": "id2", "metadata": {"user_id": "bob", "category": "work"}},
|
||||
{"key": "id3", "metadata": {"user_id": "alice", "category": "home"}},
|
||||
]
|
||||
}
|
||||
]
|
||||
store = S3Vectors(
|
||||
vector_bucket_name=BUCKET_NAME,
|
||||
collection_name=INDEX_NAME,
|
||||
embedding_model_dims=EMBEDDING_DIMS,
|
||||
)
|
||||
|
||||
[results] = store.list(filters={"user_id": "alice", "category": "work"})
|
||||
|
||||
assert [result.id for result in results] == ["id1"]
|
||||
|
||||
Reference in New Issue
Block a user