fix: fix shard key selector usage in update payload methods

Co-authored-by: 2sumtech <2sumtech@gmail.com>
This commit is contained in:
George Panchuk
2026-08-20 20:12:12 +07:00
co-authored by 2sumtech
parent f003e6c525
commit 6523d6917c
3 changed files with 234 additions and 34 deletions
+47 -17
View File
@@ -1108,7 +1108,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
) -> types.UpdateResult:
if self._prefer_grpc:
(points_selector, opt_shard_key_selector) = self._try_argument_to_grpc_selector(points)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
if isinstance(shard_key_selector, get_args_subscribed(models.ShardKeySelector)):
@@ -1130,7 +1132,7 @@ class AsyncQdrantRemote(AsyncQdrantBase):
assert grpc_result is not None, "Delete vectors returned None result"
return GrpcToRest.convert_update_result(grpc_result)
else:
(_points, _filter) = self._try_argument_to_rest_points_and_filter(points)
(_points, _filter, _shard_key) = self._try_argument_to_rest_points_and_filter(points)
return (
await self.openapi_client.points_api.delete_vectors(
collection_name=collection_name,
@@ -1142,7 +1144,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
vector=vectors,
points=_points,
filter=_filter,
shard_key=shard_key_selector,
shard_key=shard_key_selector
if shard_key_selector is not None
else _shard_key,
),
)
).result
@@ -1260,7 +1264,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
points_selector.shard_key = shard_key_selector
elif isinstance(points, get_args(models.PointsSelector)):
points_selector = points
points_selector.shard_key = shard_key_selector
points_selector.shard_key = (
shard_key_selector if shard_key_selector is not None else points_selector.shard_key
)
elif isinstance(points, models.Filter):
points_selector = construct(
models.FilterSelector, filter=points, shard_key=shard_key_selector
@@ -1290,9 +1296,12 @@ class AsyncQdrantRemote(AsyncQdrantBase):
@classmethod
def _try_argument_to_rest_points_and_filter(
cls, points: types.PointsSelector
) -> tuple[list[models.ExtendedPointId] | None, models.Filter | None]:
) -> tuple[
list[models.ExtendedPointId] | None, models.Filter | None, models.ShardKeySelector | None
]:
_points = None
_filter = None
_shard_key = None
if isinstance(points, list):
_points = [
GrpcToRest.convert_point_id(idx) if isinstance(idx, grpc.PointId) else idx
@@ -1306,15 +1315,17 @@ class AsyncQdrantRemote(AsyncQdrantBase):
_filter = selector.filter
elif isinstance(points, models.PointIdsList):
_points = points.points
_shard_key = points.shard_key
elif isinstance(points, models.FilterSelector):
_filter = points.filter
_shard_key = points.shard_key
elif isinstance(points, models.Filter):
_filter = points
elif isinstance(points, grpc.Filter):
_filter = GrpcToRest.convert_filter(points)
else:
raise ValueError(f"Unsupported points selector type: {type(points)}")
return (_points, _filter)
return (_points, _filter, _shard_key)
async def delete(
self,
@@ -1330,7 +1341,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
(points_selector, opt_shard_key_selector) = self._try_argument_to_grpc_selector(
points_selector
)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
if isinstance(shard_key_selector, get_args_subscribed(models.ShardKeySelector)):
@@ -1380,7 +1393,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
) -> types.UpdateResult:
if self._prefer_grpc:
(points_selector, opt_shard_key_selector) = self._try_argument_to_grpc_selector(points)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
if isinstance(shard_key_selector, get_args_subscribed(models.ShardKeySelector)):
@@ -1403,7 +1418,7 @@ class AsyncQdrantRemote(AsyncQdrantBase):
).result
)
else:
(_points, _filter) = self._try_argument_to_rest_points_and_filter(points)
(_points, _filter, _shard_key) = self._try_argument_to_rest_points_and_filter(points)
result: types.UpdateResult | None = (
await self.openapi_client.points_api.set_payload(
collection_name=collection_name,
@@ -1414,7 +1429,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
payload=payload,
points=_points,
filter=_filter,
shard_key=shard_key_selector,
shard_key=shard_key_selector
if shard_key_selector is not None
else _shard_key,
key=key,
),
)
@@ -1435,7 +1452,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
) -> types.UpdateResult:
if self._prefer_grpc:
(points_selector, opt_shard_key_selector) = self._try_argument_to_grpc_selector(points)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
if isinstance(shard_key_selector, get_args_subscribed(models.ShardKeySelector)):
@@ -1457,7 +1476,7 @@ class AsyncQdrantRemote(AsyncQdrantBase):
).result
)
else:
(_points, _filter) = self._try_argument_to_rest_points_and_filter(points)
(_points, _filter, _shard_key) = self._try_argument_to_rest_points_and_filter(points)
result: types.UpdateResult | None = (
await self.openapi_client.points_api.overwrite_payload(
collection_name=collection_name,
@@ -1468,7 +1487,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
payload=payload,
points=_points,
filter=_filter,
shard_key=shard_key_selector,
shard_key=shard_key_selector
if shard_key_selector is not None
else _shard_key,
),
)
).result
@@ -1488,7 +1509,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
) -> types.UpdateResult:
if self._prefer_grpc:
(points_selector, opt_shard_key_selector) = self._try_argument_to_grpc_selector(points)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
if isinstance(shard_key_selector, get_args_subscribed(models.ShardKeySelector)):
@@ -1510,7 +1533,7 @@ class AsyncQdrantRemote(AsyncQdrantBase):
).result
)
else:
(_points, _filter) = self._try_argument_to_rest_points_and_filter(points)
(_points, _filter, _shard_key) = self._try_argument_to_rest_points_and_filter(points)
result: types.UpdateResult | None = (
await self.openapi_client.points_api.delete_payload(
collection_name=collection_name,
@@ -1518,7 +1541,12 @@ class AsyncQdrantRemote(AsyncQdrantBase):
ordering=ordering,
timeout=timeout,
delete_payload=models.DeletePayload(
keys=keys, points=_points, filter=_filter, shard_key=shard_key_selector
keys=keys,
points=_points,
filter=_filter,
shard_key=shard_key_selector
if shard_key_selector is not None
else _shard_key,
),
)
).result
@@ -1539,7 +1567,9 @@ class AsyncQdrantRemote(AsyncQdrantBase):
(points_selector, opt_shard_key_selector) = self._try_argument_to_grpc_selector(
points_selector
)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
if isinstance(shard_key_selector, get_args_subscribed(models.ShardKeySelector)):
+36 -17
View File
@@ -1254,7 +1254,9 @@ class QdrantRemote(QdrantBase):
) -> types.UpdateResult:
if self._prefer_grpc:
points_selector, opt_shard_key_selector = self._try_argument_to_grpc_selector(points)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
@@ -1281,7 +1283,7 @@ class QdrantRemote(QdrantBase):
return GrpcToRest.convert_update_result(grpc_result)
else:
_points, _filter = self._try_argument_to_rest_points_and_filter(points)
_points, _filter, _shard_key = self._try_argument_to_rest_points_and_filter(points)
return self.openapi_client.points_api.delete_vectors(
collection_name=collection_name,
wait=wait,
@@ -1292,7 +1294,7 @@ class QdrantRemote(QdrantBase):
vector=vectors,
points=_points,
filter=_filter,
shard_key=shard_key_selector,
shard_key=shard_key_selector if shard_key_selector is not None else _shard_key,
),
).result
@@ -1423,7 +1425,9 @@ class QdrantRemote(QdrantBase):
points_selector.shard_key = shard_key_selector
elif isinstance(points, get_args(models.PointsSelector)):
points_selector = points
points_selector.shard_key = shard_key_selector
points_selector.shard_key = (
shard_key_selector if shard_key_selector is not None else points_selector.shard_key
)
elif isinstance(points, models.Filter):
points_selector = construct(
models.FilterSelector, filter=points, shard_key=shard_key_selector
@@ -1455,9 +1459,12 @@ class QdrantRemote(QdrantBase):
@classmethod
def _try_argument_to_rest_points_and_filter(
cls, points: types.PointsSelector
) -> tuple[list[models.ExtendedPointId] | None, models.Filter | None]:
) -> tuple[
list[models.ExtendedPointId] | None, models.Filter | None, models.ShardKeySelector | None
]:
_points = None
_filter = None
_shard_key = None
if isinstance(points, list):
_points = [
(GrpcToRest.convert_point_id(idx) if isinstance(idx, grpc.PointId) else idx)
@@ -1471,8 +1478,10 @@ class QdrantRemote(QdrantBase):
_filter = selector.filter
elif isinstance(points, models.PointIdsList):
_points = points.points
_shard_key = points.shard_key
elif isinstance(points, models.FilterSelector):
_filter = points.filter
_shard_key = points.shard_key
elif isinstance(points, models.Filter):
_filter = points
elif isinstance(points, grpc.Filter):
@@ -1480,7 +1489,7 @@ class QdrantRemote(QdrantBase):
else:
raise ValueError(f"Unsupported points selector type: {type(points)}")
return _points, _filter
return _points, _filter, _shard_key
def delete(
self,
@@ -1496,7 +1505,9 @@ class QdrantRemote(QdrantBase):
points_selector, opt_shard_key_selector = self._try_argument_to_grpc_selector(
points_selector
)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
@@ -1545,7 +1556,9 @@ class QdrantRemote(QdrantBase):
) -> types.UpdateResult:
if self._prefer_grpc:
points_selector, opt_shard_key_selector = self._try_argument_to_grpc_selector(points)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
@@ -1569,7 +1582,7 @@ class QdrantRemote(QdrantBase):
).result
)
else:
_points, _filter = self._try_argument_to_rest_points_and_filter(points)
_points, _filter, _shard_key = self._try_argument_to_rest_points_and_filter(points)
result: types.UpdateResult | None = self.openapi_client.points_api.set_payload(
collection_name=collection_name,
wait=wait,
@@ -1579,7 +1592,7 @@ class QdrantRemote(QdrantBase):
payload=payload,
points=_points,
filter=_filter,
shard_key=shard_key_selector,
shard_key=shard_key_selector if shard_key_selector is not None else _shard_key,
key=key,
),
).result
@@ -1599,7 +1612,9 @@ class QdrantRemote(QdrantBase):
) -> types.UpdateResult:
if self._prefer_grpc:
points_selector, opt_shard_key_selector = self._try_argument_to_grpc_selector(points)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
@@ -1622,7 +1637,7 @@ class QdrantRemote(QdrantBase):
).result
)
else:
_points, _filter = self._try_argument_to_rest_points_and_filter(points)
_points, _filter, _shard_key = self._try_argument_to_rest_points_and_filter(points)
result: types.UpdateResult | None = self.openapi_client.points_api.overwrite_payload(
collection_name=collection_name,
wait=wait,
@@ -1632,7 +1647,7 @@ class QdrantRemote(QdrantBase):
payload=payload,
points=_points,
filter=_filter,
shard_key=shard_key_selector,
shard_key=shard_key_selector if shard_key_selector is not None else _shard_key,
),
).result
assert result is not None, "Overwrite payload returned None"
@@ -1651,7 +1666,9 @@ class QdrantRemote(QdrantBase):
) -> types.UpdateResult:
if self._prefer_grpc:
points_selector, opt_shard_key_selector = self._try_argument_to_grpc_selector(points)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
@@ -1673,7 +1690,7 @@ class QdrantRemote(QdrantBase):
).result
)
else:
_points, _filter = self._try_argument_to_rest_points_and_filter(points)
_points, _filter, _shard_key = self._try_argument_to_rest_points_and_filter(points)
result: types.UpdateResult | None = self.openapi_client.points_api.delete_payload(
collection_name=collection_name,
wait=wait,
@@ -1683,7 +1700,7 @@ class QdrantRemote(QdrantBase):
keys=keys,
points=_points,
filter=_filter,
shard_key=shard_key_selector,
shard_key=shard_key_selector if shard_key_selector is not None else _shard_key,
),
).result
assert result is not None, "Delete payload returned None"
@@ -1703,7 +1720,9 @@ class QdrantRemote(QdrantBase):
points_selector, opt_shard_key_selector = self._try_argument_to_grpc_selector(
points_selector
)
shard_key_selector = shard_key_selector or opt_shard_key_selector
shard_key_selector = (
shard_key_selector if shard_key_selector is not None else opt_shard_key_selector
)
if isinstance(ordering, models.WriteOrdering):
ordering = RestToGrpc.convert_write_ordering(ordering)
+151
View File
@@ -1254,6 +1254,157 @@ def test_custom_sharding(prefer_grpc):
client.delete_shard_key(collection_name=COLLECTION_NAME, shard_key=dogs_shard_key)
@pytest.mark.parametrize("prefer_grpc", [False, True])
def test_shard_key_from_points_selector(prefer_grpc):
"""A shard key set inside a points selector must be used, and an explicitly passed
`shard_key_selector` must take precedence over it, even if it is falsy."""
client = QdrantClient(prefer_grpc=prefer_grpc, timeout=TIMEOUT)
if client.cluster_status().status == "disabled":
pytest.skip("Requires distributed mode")
# `0` is a valid shard key, it must not be treated as a missing one
zero_shard_key, one_shard_key = 0, 1
shard_key_ids = {zero_shard_key: [1, 2], one_shard_key: [3, 4]}
cat_payload = {"name": "Barsik"}
dog_payload = {"name": "Sharik"}
def points(shard_key):
return [
PointStruct(
id=point_id, vector={"text": np.random.rand(DIM).tolist()}, payload=cat_payload
)
for point_id in shard_key_ids[shard_key]
]
def init_collection():
if client.collection_exists(COLLECTION_NAME):
client.delete_collection(collection_name=COLLECTION_NAME)
client.create_collection(
collection_name=COLLECTION_NAME,
vectors_config={"text": VectorParams(size=DIM, distance=Distance.DOT)},
sharding_method=models.ShardingMethod.CUSTOM,
)
for shard_key in shard_key_ids:
client.create_shard_key(collection_name=COLLECTION_NAME, shard_key=shard_key)
reset_points()
# restores payloads, vectors and points changed by the previous case
def reset_points():
for shard_key in shard_key_ids:
client.upsert(
collection_name=COLLECTION_NAME,
points=points(shard_key),
shard_key_selector=shard_key,
)
def records(shard_key):
return client.scroll(
collection_name=COLLECTION_NAME,
shard_key_selector=shard_key,
with_vectors=True,
limit=10,
)[0]
def payloads(shard_key):
return [record.payload for record in records(shard_key)]
# both selectors are meant for the `zero_shard_key` shard: the ids are the ones of its
# points, the filter matches every point of whichever shard the operation reaches
def selectors(selector_shard_key):
return [
models.PointIdsList(
points=shard_key_ids[zero_shard_key], shard_key=selector_shard_key
),
models.FilterSelector(filter=models.Filter(), shard_key=selector_shard_key),
]
# shard key in the points selector, explicitly passed shard_key_selector. Every combination
# must apply the operation to the `zero_shard_key` shard, and only to it
shard_key_combinations = [
(zero_shard_key, None), # taken from the points selector
(one_shard_key, zero_shard_key), # explicit one wins over the one in the selector
(None, zero_shard_key), # explicit one is used if the selector has none
]
def cases():
for selector_shard_key, explicit_shard_key in shard_key_combinations:
for points_selector in selectors(selector_shard_key):
reset_points()
yield points_selector, explicit_shard_key
init_collection()
# region delete_payload
for points_selector, explicit_shard_key in cases():
client.delete_payload(
COLLECTION_NAME,
keys=["name"],
points=points_selector,
shard_key_selector=explicit_shard_key,
)
assert payloads(zero_shard_key) == [{}, {}]
assert payloads(one_shard_key) == [cat_payload, cat_payload]
# endregion
# region set_payload
for points_selector, explicit_shard_key in cases():
client.set_payload(
COLLECTION_NAME,
payload={"age": 3},
points=points_selector,
shard_key_selector=explicit_shard_key,
)
assert payloads(zero_shard_key) == [{**cat_payload, "age": 3}] * 2
assert payloads(one_shard_key) == [cat_payload, cat_payload]
# endregion
# region overwrite_payload
for points_selector, explicit_shard_key in cases():
client.overwrite_payload(
COLLECTION_NAME,
payload=dog_payload,
points=points_selector,
shard_key_selector=explicit_shard_key,
)
assert payloads(zero_shard_key) == [dog_payload, dog_payload]
assert payloads(one_shard_key) == [cat_payload, cat_payload]
# endregion
# region clear_payload
for points_selector, explicit_shard_key in cases():
client.clear_payload(
COLLECTION_NAME,
points_selector=points_selector,
shard_key_selector=explicit_shard_key,
)
assert payloads(zero_shard_key) == [{}, {}]
assert payloads(one_shard_key) == [cat_payload, cat_payload]
# endregion
# region delete_vectors
for points_selector, explicit_shard_key in cases():
client.delete_vectors(
COLLECTION_NAME,
vectors=["text"],
points=points_selector,
shard_key_selector=explicit_shard_key,
)
assert [record.vector for record in records(zero_shard_key)] == [{}, {}]
assert all(record.vector for record in records(one_shard_key))
# endregion
# region delete
for points_selector, explicit_shard_key in cases():
client.delete(
COLLECTION_NAME,
points_selector=points_selector,
shard_key_selector=explicit_shard_key,
)
assert records(zero_shard_key) == []
assert [record.id for record in records(one_shard_key)] == shard_key_ids[one_shard_key]
# endregion
@pytest.mark.parametrize("prefer_grpc", [False, True])
def test_sparse_vectors(prefer_grpc):
client = QdrantClient(prefer_grpc=prefer_grpc, timeout=TIMEOUT)