fix(vector_stores): handle vector=None in Milvus and Qdrant update methods (#4568)

Co-authored-by: utkarsh240799 <utkarsh240799@users.noreply.github.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Utkarsh
2026-03-27 21:50:54 +05:30
committed by GitHub
parent 68cf4e118d
commit e87240d4d9
4 changed files with 134 additions and 2 deletions
+11
View File
@@ -181,6 +181,17 @@ class MilvusDB(VectorStoreBase):
vector (List[float], optional): Updated vector.
payload (Dict, optional): Updated payload.
"""
if vector is None or payload is None:
existing = self.client.get(collection_name=self.collection_name, ids=vector_id)
if not existing:
raise ValueError(f"Vector with id {vector_id} not found in collection {self.collection_name}")
if vector is None:
vector = existing[0].get("vectors")
if vector is None:
raise ValueError(f"Existing record {vector_id} has no vector data")
if payload is None:
payload = existing[0].get("metadata")
schema = {"id": vector_id, "vectors": vector, "metadata": payload}
self.client.upsert(collection_name=self.collection_name, data=schema)
+16 -2
View File
@@ -12,6 +12,7 @@ from qdrant_client.models import (
MatchValue,
PointIdsList,
PointStruct,
PointVectors,
Range,
VectorParams,
)
@@ -326,8 +327,21 @@ class Qdrant(VectorStoreBase):
vector (list, optional): Updated vector. Defaults to None.
payload (dict, optional): Updated payload. Defaults to None.
"""
point = PointStruct(id=vector_id, vector=vector, payload=payload)
self.client.upsert(collection_name=self.collection_name, points=[point])
if vector is not None and payload is not None:
point = PointStruct(id=vector_id, vector=vector, payload=payload)
self.client.upsert(collection_name=self.collection_name, points=[point])
else:
if payload is not None:
self.client.set_payload(
collection_name=self.collection_name,
payload=payload,
points=[vector_id],
)
if vector is not None:
self.client.update_vectors(
collection_name=self.collection_name,
points=[PointVectors(id=vector_id, vector=vector)],
)
def get(self, vector_id: int) -> dict:
"""
+69
View File
@@ -233,6 +233,75 @@ class TestMilvusDB:
assert parsed[1].id == "mem2"
assert parsed[1].score == 0.85
def test_update_with_none_vector_fetches_existing(self, milvus_db, mock_milvus_client):
"""Test that update with vector=None fetches the existing vector (fixes #3708)."""
vector_id = "test_id"
existing_vector = [0.5] * 1536
payload = {"user_id": "alice", "data": "Updated memory"}
mock_milvus_client.get.return_value = [
{"id": vector_id, "vectors": existing_vector, "metadata": {"user_id": "alice"}}
]
milvus_db.update(vector_id=vector_id, vector=None, payload=payload)
mock_milvus_client.get.assert_called_once_with(
collection_name="test_collection", ids=vector_id
)
call_args = mock_milvus_client.upsert.call_args
assert call_args[1]['data']['vectors'] == existing_vector
assert call_args[1]['data']['metadata'] == payload
def test_update_with_none_payload_fetches_existing(self, milvus_db, mock_milvus_client):
"""Test that update with payload=None fetches the existing metadata."""
vector_id = "test_id"
vector = [0.1] * 1536
existing_metadata = {"user_id": "alice", "data": "Original"}
mock_milvus_client.get.return_value = [
{"id": vector_id, "vectors": [0.5] * 1536, "metadata": existing_metadata}
]
milvus_db.update(vector_id=vector_id, vector=vector, payload=None)
call_args = mock_milvus_client.upsert.call_args
assert call_args[1]['data']['vectors'] == vector
assert call_args[1]['data']['metadata'] == existing_metadata
def test_update_with_both_none_fetches_existing(self, milvus_db, mock_milvus_client):
"""Test that update with both vector=None and payload=None fetches existing data."""
vector_id = "test_id"
existing_vector = [0.5] * 1536
existing_metadata = {"user_id": "alice"}
mock_milvus_client.get.return_value = [
{"id": vector_id, "vectors": existing_vector, "metadata": existing_metadata}
]
milvus_db.update(vector_id=vector_id, vector=None, payload=None)
# Should only call get once even though both are None
assert mock_milvus_client.get.call_count == 1
call_args = mock_milvus_client.upsert.call_args
assert call_args[1]['data']['vectors'] == existing_vector
assert call_args[1]['data']['metadata'] == existing_metadata
def test_update_with_none_vector_raises_on_missing_record(self, milvus_db, mock_milvus_client):
"""Test that update raises ValueError when the record doesn't exist."""
mock_milvus_client.get.return_value = []
with pytest.raises(ValueError, match="not found"):
milvus_db.update(vector_id="nonexistent", vector=None, payload={"data": "test"})
def test_update_with_none_vector_raises_on_missing_vector_data(self, milvus_db, mock_milvus_client):
"""Test that update raises ValueError when existing record has no vector."""
mock_milvus_client.get.return_value = [
{"id": "test_id", "vectors": None, "metadata": {"user_id": "alice"}}
]
with pytest.raises(ValueError, match="no vector data"):
milvus_db.update(vector_id="test_id", vector=None, payload={"data": "test"})
def test_collection_already_exists(self, mock_milvus_client):
"""Test that existing collection is not recreated."""
mock_milvus_client.has_collection.return_value = True
+38
View File
@@ -15,6 +15,7 @@ from qdrant_client.models import (
MatchValue,
PointIdsList,
PointStruct,
PointVectors,
Range,
VectorParams,
)
@@ -232,6 +233,43 @@ class TestQdrant(unittest.TestCase):
self.assertEqual(point.vector, updated_vector)
self.assertEqual(point.payload, updated_payload)
def test_update_with_none_vector_uses_set_payload(self):
"""Test that update with vector=None uses set_payload instead of upsert (fixes #3708)."""
vector_id = str(uuid.uuid4())
updated_payload = {"key": "updated_value"}
self.qdrant.update(vector_id=vector_id, vector=None, payload=updated_payload)
self.client_mock.upsert.assert_not_called()
self.client_mock.set_payload.assert_called_once_with(
collection_name="test_collection",
payload=updated_payload,
points=[vector_id],
)
def test_update_with_none_payload_uses_update_vectors(self):
"""Test that update with payload=None uses update_vectors instead of upsert."""
vector_id = str(uuid.uuid4())
updated_vector = [0.2, 0.3]
self.qdrant.update(vector_id=vector_id, vector=updated_vector, payload=None)
self.client_mock.upsert.assert_not_called()
self.client_mock.update_vectors.assert_called_once_with(
collection_name="test_collection",
points=[PointVectors(id=vector_id, vector=updated_vector)],
)
def test_update_with_both_none_is_noop(self):
"""Test that update with both vector=None and payload=None is a no-op."""
vector_id = str(uuid.uuid4())
self.qdrant.update(vector_id=vector_id, vector=None, payload=None)
self.client_mock.upsert.assert_not_called()
self.client_mock.set_payload.assert_not_called()
self.client_mock.update_vectors.assert_not_called()
def test_get(self):
vector_id = str(uuid.uuid4())
self.client_mock.retrieve.return_value = [{"id": vector_id, "payload": {"key": "value"}}]