mirror of
https://github.com/qdrant/qdrant-client.git
synced 2026-09-21 13:37:55 -05:00
fix: fix shard key selector usage in update payload methods
Co-authored-by: 2sumtech <2sumtech@gmail.com>
This commit is contained in:
co-authored by
2sumtech
parent
f003e6c525
commit
6523d6917c
@@ -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)):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user