diff --git a/qdrant_client/async_qdrant_remote.py b/qdrant_client/async_qdrant_remote.py index ea139e18..99c1ad15 100644 --- a/qdrant_client/async_qdrant_remote.py +++ b/qdrant_client/async_qdrant_remote.py @@ -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)): diff --git a/qdrant_client/qdrant_remote.py b/qdrant_client/qdrant_remote.py index 96789697..3ce6b0d9 100644 --- a/qdrant_client/qdrant_remote.py +++ b/qdrant_client/qdrant_remote.py @@ -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) diff --git a/tests/test_qdrant_client.py b/tests/test_qdrant_client.py index db765120..182978fb 100644 --- a/tests/test_qdrant_client.py +++ b/tests/test_qdrant_client.py @@ -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)