new: add py.typed, validate code with mypy (#91)

This commit is contained in:
George
2023-01-21 16:32:32 +04:00
committed by GitHub
parent 71f73c271f
commit 3e68b3786f
8 changed files with 208 additions and 150 deletions

4
mypy.ini Normal file
View File

@@ -0,0 +1,4 @@
[mypy]
ignore_missing_imports = True
follow_imports = skip
exclude = qdrant_client/grpc|qdrant_client/http|tests

View File

@@ -1,3 +1,10 @@
import sys
if sys.version_info >= (3, 10):
from typing import TypeAlias
else:
from typing_extensions import TypeAlias
from typing import Union, List
from qdrant_client import grpc as grpc
@@ -23,19 +30,19 @@ AliasOperations = Union[
rest.DeleteAliasOperation,
grpc.AliasOperations
]
Payload = rest.Payload
Payload: TypeAlias = rest.Payload
ScoredPoint = rest.ScoredPoint
UpdateResult = rest.UpdateResult
Record = rest.Record
CollectionsResponse = rest.CollectionsResponse
CollectionInfo = rest.CollectionInfo
CountResult = rest.CountResult
SnapshotDescription = rest.SnapshotDescription
NamedVector = rest.NamedVector
VectorParams = rest.VectorParams
LocksOption = rest.LocksOption
SnapshotPriority = rest.SnapshotPriority
ScoredPoint: TypeAlias = rest.ScoredPoint
UpdateResult: TypeAlias = rest.UpdateResult
Record: TypeAlias = rest.Record
CollectionsResponse: TypeAlias = rest.CollectionsResponse
CollectionInfo: TypeAlias = rest.CollectionInfo
CountResult: TypeAlias = rest.CountResult
SnapshotDescription: TypeAlias = rest.SnapshotDescription
NamedVector: TypeAlias = rest.NamedVector
VectorParams: TypeAlias = rest.VectorParams
LocksOption: TypeAlias = rest.LocksOption
SnapshotPriority: TypeAlias = rest.SnapshotPriority
SearchRequest = Union[rest.SearchRequest, grpc.SearchPoints]
RecommendRequest = Union[rest.RecommendRequest, grpc.RecommendPoints]

View File

@@ -5,7 +5,7 @@ from google.protobuf.timestamp_pb2 import Timestamp
from google.protobuf.json_format import MessageToDict
try:
from google.protobuf.pyext._message import MessageMapContainer
from google.protobuf.pyext._message import MessageMapContainer # type: ignore
except ImportError:
pass
@@ -1135,13 +1135,10 @@ class RestToGrpc:
@classmethod
def convert_batch_vector_struct(cls, model: rest.BatchVectorStruct, num_records: int) -> List[grpc.Vectors]:
"""
"""
if isinstance(model, list):
return [cls.convert_vector_struct(item) for item in model]
elif isinstance(model, dict):
result = [{} for _ in range(num_records)]
result: List[Dict] = [{} for _ in range(num_records)]
for key, val in model.items():
for i, item in enumerate(val):
result[i][key] = item
@@ -1155,6 +1152,8 @@ class RestToGrpc:
return model, None
elif isinstance(model, rest.NamedVector):
return model.vector, model.name
else:
raise ValueError(f"invalid NamedVectorStruct model: {model}")
@classmethod
def convert_search_request(cls, model: rest.SearchRequest, collection_name: str) -> grpc.SearchPoints:

View File

@@ -1,6 +1,7 @@
import os
from enum import Enum
from multiprocessing import Queue, Value, get_context
from multiprocessing import Queue, get_context
from multiprocessing.sharedctypes import Synchronized as BaseValue
from multiprocessing.context import BaseContext
from multiprocessing.process import BaseProcess
from queue import Empty
@@ -32,7 +33,7 @@ class Worker:
def _worker(worker_class: Type[Worker],
input_queue: Queue,
output_queue: Queue,
num_active_workers: Value,
num_active_workers: BaseValue,
worker_id: int,
kwargs=None) -> None:
"""
@@ -90,7 +91,7 @@ class ParallelWorkerPool:
self.processes: List[BaseProcess] = []
self.queue_size = self.num_workers * max_internal_batch_size
self.num_active_workers: Optional[Value] = None
self.num_active_workers: Optional[BaseValue] = None
def start(self, **kwargs):
self.input_queue = self.ctx.Queue(self.queue_size)
@@ -115,6 +116,10 @@ class ParallelWorkerPool:
def unordered_map(self, stream: Iterable[Any], *args, **kwargs) -> Iterable[Any]:
try:
self.start(**kwargs)
assert self.input_queue is not None, "Input queue was not initialized"
assert self.output_queue is not None, "Output queue was not initialized"
pushed = 0
read = 0
for item in stream:
@@ -152,6 +157,8 @@ class ParallelWorkerPool:
yield out_item
read += 1
finally:
assert self.input_queue is not None, "Input queue is None"
assert self.output_queue is not None, "Output queue is None"
self.input_queue.close()
self.output_queue.close()

1
qdrant_client/py.typed Normal file
View File

@@ -0,0 +1 @@
partial

View File

@@ -1,4 +1,3 @@
import asyncio
import warnings
from multiprocessing import get_all_start_methods
from typing import Optional, Iterable, List, Union, Tuple, Type, Dict
@@ -12,7 +11,7 @@ from qdrant_client.connection import get_channel
from qdrant_client.conversions import common_types as types
from qdrant_client.conversions.conversion import RestToGrpc, GrpcToRest
from qdrant_client.http import SyncApis
from qdrant_client.http import models as rest
from qdrant_client.http import models as rest_models
from qdrant_client.parallel_processor import ParallelWorkerPool
from qdrant_client.uploader.grpc_uploader import GrpcBatchUploader
from qdrant_client.uploader.rest_uploader import RestBatchUploader
@@ -60,14 +59,14 @@ class QdrantClient:
"""
def __init__(self,
host="localhost",
port=6333,
grpc_port=6334,
prefer_grpc=False,
https=None,
api_key=None,
prefix=None,
timeout=None,
host: str = "localhost",
port: int = 6333,
grpc_port: int = 6334,
prefer_grpc: bool = False,
https: Optional[bool] = None,
api_key: Optional[str] = None,
prefix: Optional[str] = None,
timeout: Optional[float] = None,
**kwargs):
self._prefer_grpc = prefer_grpc
self._grpc_port = grpc_port
@@ -116,7 +115,7 @@ class QdrantClient:
if self._timeout is not None:
self._rest_args['timeout'] = self._timeout
self.openapi_client = SyncApis(host=self.rest_uri, **self._rest_args)
self.openapi_client: SyncApis = SyncApis(host=self.rest_uri, **self._rest_args)
self._grpc_channel = None
self._grpc_points_client: Optional[grpc.PointsStub] = None
@@ -191,27 +190,28 @@ class QdrantClient:
"""
if self._prefer_grpc:
requests = [
RestToGrpc.convert_search_request(r, collection_name) if isinstance(r, rest.SearchRequest) else r
RestToGrpc.convert_search_request(r, collection_name) if isinstance(r, rest_models.SearchRequest) else r
for r in requests
]
res: grpc.SearchBatchResponse = self._grpc_points_client.SearchBatch(
grpc_res: grpc.SearchBatchResponse = self.grpc_points.SearchBatch(
grpc.SearchBatchPoints(
collection_name=collection_name,
search_points=requests,
), timeout=self._timeout
)
return [[GrpcToRest.convert_scored_point(hit) for hit in r.result] for r in res.result]
return [[GrpcToRest.convert_scored_point(hit) for hit in r.result] for r in grpc_res.result]
else:
requests = [
GrpcToRest.convert_search_points(r) if isinstance(r, grpc.SearchPoints) else r
for r in requests
]
res: List[List[rest.ScoredPoint]] = self.http.points_api.search_batch_points(
http_res: List[List[rest_models.ScoredPoint]] = self.http.points_api.search_batch_points(
collection_name=collection_name,
search_request_batch=rest.SearchRequestBatch(searches=requests)
search_request_batch=rest_models.SearchRequestBatch(searches=requests)
).result
return res
return http_res
def search(self,
collection_name: str,
@@ -224,7 +224,7 @@ class QdrantClient:
with_vectors: Union[bool, List[str]] = False,
score_threshold: Optional[float] = None,
append_payload=True,
top: int = None,
top: Optional[int] = None,
) -> List[types.ScoredPoint]:
"""Search for closest vectors in collection taking into account filtering conditions
@@ -305,19 +305,19 @@ class QdrantClient:
vector_name = query_vector[0]
vector = query_vector[1]
else:
vector = query_vector
vector = list(query_vector)
if isinstance(query_filter, rest.Filter):
if isinstance(query_filter, rest_models.Filter):
query_filter = RestToGrpc.convert_filter(model=query_filter)
if isinstance(search_params, rest.SearchParams):
if isinstance(search_params, rest_models.SearchParams):
search_params = RestToGrpc.convert_search_params(search_params)
if isinstance(with_payload, (
bool,
list,
rest.PayloadSelectorInclude,
rest.PayloadSelectorExclude
rest_models.PayloadSelectorInclude,
rest_models.PayloadSelectorExclude
)):
with_payload = RestToGrpc.convert_with_payload_interface(with_payload)
@@ -327,7 +327,7 @@ class QdrantClient:
)):
with_vectors = RestToGrpc.convert_with_vectors(with_vectors)
res: grpc.SearchResponse = self._grpc_points_client.Search(grpc.SearchPoints(
res: grpc.SearchResponse = self.grpc_points.Search(grpc.SearchPoints(
collection_name=collection_name,
vector=vector,
vector_name=vector_name,
@@ -357,7 +357,7 @@ class QdrantClient:
search_result = self.http.points_api.search_points(
collection_name=collection_name,
search_request=rest.SearchRequest(
search_request=rest_models.SearchRequest(
vector=query_vector,
filter=query_filter,
limit=limit,
@@ -387,10 +387,10 @@ class QdrantClient:
"""
if self._prefer_grpc:
requests = [
RestToGrpc.convert_recommend_request(r, collection_name) if isinstance(r, rest.RecommendRequest) else r
RestToGrpc.convert_recommend_request(r, collection_name) if isinstance(r, rest_models.RecommendRequest) else r
for r in requests
]
res: grpc.SearchBatchResponse = self._grpc_points_client.RecommendBatch(
grpc_res: grpc.SearchBatchResponse = self.grpc_points.RecommendBatch(
grpc.RecommendBatchPoints(
collection_name=collection_name,
recommend_points=requests,
@@ -398,23 +398,23 @@ class QdrantClient:
timeout=self._timeout
)
return [[GrpcToRest.convert_scored_point(hit) for hit in r.result] for r in res.result]
return [[GrpcToRest.convert_scored_point(hit) for hit in r.result] for r in grpc_res.result]
else:
requests = [
GrpcToRest.convert_recommend_points(r) if isinstance(r, grpc.RecommendPoints) else r
for r in requests
]
res: List[List[rest.ScoredPoint]] = self.http.points_api.recommend_batch_points(
http_res: List[List[rest_models.ScoredPoint]] = self.http.points_api.recommend_batch_points(
collection_name=collection_name,
recommend_request_batch=rest.RecommendRequestBatch(searches=requests)
recommend_request_batch=rest_models.RecommendRequestBatch(searches=requests)
).result
return res
return http_res
def recommend(
self,
collection_name: str,
positive: List[types.PointId],
negative: List[types.PointId] = None,
negative: Optional[List[types.PointId]] = None,
query_filter: Optional[types.Filter] = None,
search_params: Optional[types.SearchParams] = None,
limit: int = 10,
@@ -424,7 +424,7 @@ class QdrantClient:
score_threshold: Optional[float] = None,
using: Optional[str] = None,
lookup_from: Optional[types.LookupLocation] = None,
top: int = None,
top: Optional[int] = None,
) -> List[types.ScoredPoint]:
"""Recommend points: search for similar points based on already stored in Qdrant examples.
@@ -497,17 +497,17 @@ class QdrantClient:
for point_id in negative
]
if isinstance(query_filter, rest.Filter):
if isinstance(query_filter, rest_models.Filter):
query_filter = RestToGrpc.convert_filter(model=query_filter)
if isinstance(search_params, rest.SearchParams):
if isinstance(search_params, rest_models.SearchParams):
search_params = RestToGrpc.convert_search_params(search_params)
if isinstance(with_payload, (
bool,
list,
rest.PayloadSelectorInclude,
rest.PayloadSelectorExclude
rest_models.PayloadSelectorInclude,
rest_models.PayloadSelectorExclude
)):
with_payload = RestToGrpc.convert_with_payload_interface(with_payload)
@@ -517,10 +517,10 @@ class QdrantClient:
)):
with_vectors = RestToGrpc.convert_with_vectors(with_vectors)
if isinstance(lookup_from, rest.LookupLocation):
if isinstance(lookup_from, rest_models.LookupLocation):
lookup_from = RestToGrpc.convert_lookup_location(lookup_from)
res: grpc.SearchResponse = self._grpc_points_client.Recommend(grpc.RecommendPoints(
res: grpc.SearchResponse = self.grpc_points.Recommend(grpc.RecommendPoints(
collection_name=collection_name,
positive=positive,
negative=negative,
@@ -559,9 +559,9 @@ class QdrantClient:
if isinstance(lookup_from, grpc.LookupLocation):
lookup_from = GrpcToRest.convert_lookup_location(lookup_from)
return self.openapi_client.points_api.recommend_points(
result = self.openapi_client.points_api.recommend_points(
collection_name=collection_name,
recommend_request=rest.RecommendRequest(
recommend_request=rest_models.RecommendRequest(
filter=query_filter,
negative=negative,
params=search_params,
@@ -575,6 +575,8 @@ class QdrantClient:
using=using
)
).result
assert result is not None, "Recommend points API returned None"
return result
def scroll(
self,
@@ -615,14 +617,14 @@ class QdrantClient:
if isinstance(offset, (int, str)):
offset = RestToGrpc.convert_extended_point_id(offset)
if isinstance(scroll_filter, rest.Filter):
if isinstance(scroll_filter, rest_models.Filter):
scroll_filter = RestToGrpc.convert_filter(model=scroll_filter)
if isinstance(with_payload, (
bool,
list,
rest.PayloadSelectorInclude,
rest.PayloadSelectorExclude
rest_models.PayloadSelectorInclude,
rest_models.PayloadSelectorExclude
)):
with_payload = RestToGrpc.convert_with_payload_interface(with_payload)
@@ -632,7 +634,7 @@ class QdrantClient:
)):
with_vectors = RestToGrpc.convert_with_vectors(with_vectors)
res: grpc.ScrollResponse = self._grpc_points_client.Scroll(grpc.ScrollPoints(
res: grpc.ScrollResponse = self.grpc_points.Scroll(grpc.ScrollPoints(
collection_name=collection_name,
filter=scroll_filter,
offset=offset,
@@ -655,9 +657,9 @@ class QdrantClient:
if isinstance(with_payload, grpc.WithPayloadSelector):
with_payload = GrpcToRest.convert_with_payload_selector(with_payload)
scroll_result: rest.ScrollResult = self.openapi_client.points_api.scroll_points(
scroll_result: Optional[rest_models.ScrollResult] = self.openapi_client.points_api.scroll_points(
collection_name=collection_name,
scroll_request=rest.ScrollRequest(
scroll_request=rest_models.ScrollRequest(
filter=scroll_filter,
limit=limit,
offset=offset,
@@ -665,6 +667,7 @@ class QdrantClient:
with_vector=with_vectors
)
).result
assert scroll_result is not None, "Scroll points API returned None result"
return scroll_result.points, scroll_result.next_page_offset
@@ -693,12 +696,12 @@ class QdrantClient:
count_result = self.openapi_client.points_api.count_points(
collection_name=collection_name,
count_request=rest.CountRequest(
count_request=rest_models.CountRequest(
filter=count_filter,
exact=exact
)
).result
assert count_result is not None, "Count points returned None result"
return count_result
def upsert(
@@ -723,9 +726,11 @@ class QdrantClient:
Operation result
"""
if self._prefer_grpc:
if isinstance(points, rest.Batch):
vectors_batch: List[grpc.Vectors] = RestToGrpc.convert_batch_vector_struct(points.vectors,
len(points.ids))
if isinstance(points, rest_models.Batch):
vectors_batch: List[grpc.Vectors] = RestToGrpc.convert_batch_vector_struct(
points.vectors,
len(points.ids)
)
points = [
grpc.PointStruct(
id=RestToGrpc.convert_extended_point_id(points.ids[idx]),
@@ -737,15 +742,21 @@ class QdrantClient:
if isinstance(points, list):
points = [
RestToGrpc.convert_point_struct(point)
if isinstance(point, rest.PointStruct) else point
if isinstance(point, rest_models.PointStruct) else point
for point in points
]
return GrpcToRest.convert_update_result(self._grpc_points_client.Upsert(grpc.UpsertPoints(
collection_name=collection_name,
wait=wait,
points=points
), timeout=self._timeout).result)
grpc_result = self.grpc_points.Upsert(
grpc.UpsertPoints(
collection_name=collection_name,
wait=wait,
points=points
),
timeout=self._timeout
).result
assert grpc_result is not None, "Upsert returned None result"
return GrpcToRest.convert_update_result(grpc_result)
else:
if isinstance(points, list):
points = [
@@ -754,16 +765,18 @@ class QdrantClient:
for point in points
]
points = rest.PointsList(points=points)
points = rest_models.PointsList(points=points)
if isinstance(points, rest.Batch):
points = rest.PointsBatch(batch=points)
if isinstance(points, rest_models.Batch):
points = rest_models.PointsBatch(batch=points)
return self.openapi_client.points_api.upsert_points(
http_result = self.openapi_client.points_api.upsert_points(
collection_name=collection_name,
wait=wait,
point_insert_operations=points
).result
assert http_result is not None, "Upsert returned None result"
return http_result
def retrieve(
self,
@@ -796,8 +809,8 @@ class QdrantClient:
if isinstance(with_payload, (
bool,
list,
rest.PayloadSelectorInclude,
rest.PayloadSelectorExclude
rest_models.PayloadSelectorInclude,
rest_models.PayloadSelectorExclude
)):
with_payload = RestToGrpc.convert_with_payload_interface(with_payload)
@@ -812,12 +825,17 @@ class QdrantClient:
)):
with_vectors = RestToGrpc.convert_with_vectors(with_vectors)
result = self._grpc_points_client.Get(grpc.GetPoints(
collection_name=collection_name,
ids=ids,
with_payload=with_payload,
with_vectors=with_vectors
), timeout=self._timeout).result
result = self.grpc_points.Get(
grpc.GetPoints(
collection_name=collection_name,
ids=ids,
with_payload=with_payload,
with_vectors=with_vectors
),
timeout=self._timeout
).result
assert result is not None, "Retrieve returned None result"
return [
GrpcToRest.convert_retrieved_point(record)
@@ -833,14 +851,16 @@ class QdrantClient:
for idx in ids
]
return self.openapi_client.points_api.get_points(
http_result = self.openapi_client.points_api.get_points(
collection_name=collection_name,
point_request=rest.PointRequest(
point_request=rest_models.PointRequest(
ids=ids,
with_payload=with_payload,
with_vector=with_vectors
)
).result
assert http_result is not None, "Retrieve API returned None result"
return http_result
@classmethod
def _try_argument_to_grpc_selector(
@@ -855,10 +875,10 @@ class QdrantClient:
]))
elif isinstance(points, grpc.PointsSelector):
points_selector = points
elif isinstance(points, (rest.PointIdsList, rest.FilterSelector)):
elif isinstance(points, (rest_models.PointIdsList, rest_models.FilterSelector)):
points_selector = RestToGrpc.convert_points_selector(points)
elif isinstance(points, rest.Filter):
points_selector = RestToGrpc.convert_points_selector(rest.FilterSelector.construct(filter=points))
elif isinstance(points, rest_models.Filter):
points_selector = RestToGrpc.convert_points_selector(rest_models.FilterSelector.construct(filter=points))
elif isinstance(points, grpc.Filter):
points_selector = grpc.PointsSelector(
filter=points
@@ -871,27 +891,30 @@ class QdrantClient:
def _try_argument_to_rest_selector(
cls,
points: types.PointsSelector
) -> 'rest.PointsSelector':
) -> rest_models.PointsSelector:
if isinstance(points, list):
_points = [
GrpcToRest.convert_point_id(idx) if isinstance(idx, grpc.PointId) else idx
for idx in points
]
points_selector = rest.PointIdsList.construct(points=_points)
points_selector = rest_models.PointIdsList.construct(points=_points)
elif isinstance(points, grpc.PointsSelector):
points_selector = GrpcToRest.convert_points_selector(points)
elif isinstance(points, (rest.PointIdsList, rest.FilterSelector)):
elif isinstance(points, (rest_models.PointIdsList, rest_models.FilterSelector)):
points_selector = points
elif isinstance(points, rest.Filter):
points_selector = rest.FilterSelector.construct(filter=points)
elif isinstance(points, rest_models.Filter):
points_selector = rest_models.FilterSelector.construct(filter=points)
elif isinstance(points, grpc.Filter):
points_selector = rest.FilterSelector.construct(filter=GrpcToRest.convert_filter(points))
points_selector = rest_models.FilterSelector.construct(filter=GrpcToRest.convert_filter(points))
else:
raise ValueError(f"Unsupported points selector type: {type(points)}")
return points_selector
@classmethod
def _points_selector_to_points_list(cls, points_selector: grpc.PointsSelector) -> List[grpc.PointId]:
def _points_selector_to_points_list(
cls,
points_selector: grpc.PointsSelector
) -> List[grpc.PointId]:
name = points_selector.WhichOneof("points_selector_one_of")
val = getattr(points_selector, name)
@@ -903,7 +926,7 @@ class QdrantClient:
def _try_argument_to_rest_points_and_filter(
cls,
points: types.PointsSelector
) -> 'Tuple[Optional[List[rest.ExtendedPointId]], Optional[rest.Filter]]':
) -> Tuple[Optional[List[rest_models.ExtendedPointId]], Optional[rest_models.Filter]]:
_points = None
_filter = None
if isinstance(points, list):
@@ -913,15 +936,15 @@ class QdrantClient:
]
elif isinstance(points, grpc.PointsSelector):
selector = GrpcToRest.convert_points_selector(points)
if isinstance(selector, rest.PointIdsList):
if isinstance(selector, rest_models.PointIdsList):
_points = selector.points
elif isinstance(selector, rest.FilterSelector):
elif isinstance(selector, rest_models.FilterSelector):
_filter = selector.filter
elif isinstance(points, rest.PointIdsList):
elif isinstance(points, rest_models.PointIdsList):
_points = points.points
elif isinstance(points, rest.FilterSelector):
elif isinstance(points, rest_models.FilterSelector):
_filter = points.filter
elif isinstance(points, rest.Filter):
elif isinstance(points, rest_models.Filter):
_filter = points
elif isinstance(points, grpc.Filter):
_filter = GrpcToRest.convert_filter(points)
@@ -956,7 +979,7 @@ class QdrantClient:
points_selector = self._try_argument_to_grpc_selector(points_selector)
return GrpcToRest.convert_update_result(
self._grpc_points_client.Delete(grpc.DeletePoints(
self.grpc_points.Delete(grpc.DeletePoints(
collection_name=collection_name,
wait=wait,
points=points_selector
@@ -1013,7 +1036,7 @@ class QdrantClient:
points_selector = self._try_argument_to_grpc_selector(points)
return GrpcToRest.convert_update_result(
self._grpc_points_client.SetPayload(grpc.SetPayloadPoints(
self.grpc_points.SetPayload(grpc.SetPayloadPoints(
collection_name=collection_name,
wait=wait,
payload=RestToGrpc.convert_payload(payload),
@@ -1026,7 +1049,7 @@ class QdrantClient:
return self.openapi_client.points_api.set_payload(
collection_name=collection_name,
wait=wait,
set_payload=rest.SetPayload(
set_payload=rest_models.SetPayload(
payload=payload,
points=_points,
filter=_filter,
@@ -1078,7 +1101,7 @@ class QdrantClient:
points_selector = self._try_argument_to_grpc_selector(points)
return GrpcToRest.convert_update_result(
self._grpc_points_client.OverwritePayload(grpc.SetPayloadPoints(
self.grpc_points.OverwritePayload(grpc.SetPayloadPoints(
collection_name=collection_name,
wait=wait,
payload=RestToGrpc.convert_payload(payload),
@@ -1091,7 +1114,7 @@ class QdrantClient:
return self.openapi_client.points_api.overwrite_payload(
collection_name=collection_name,
wait=wait,
set_payload=rest.SetPayload(
set_payload=rest_models.SetPayload(
payload=payload,
points=_points,
filter=_filter,
@@ -1125,7 +1148,7 @@ class QdrantClient:
if self._prefer_grpc:
points_selector = self._try_argument_to_grpc_selector(points)
return GrpcToRest.convert_update_result(
self._grpc_points_client.DeletePayload(grpc.DeletePayloadPoints(
self.grpc_points.DeletePayload(grpc.DeletePayloadPoints(
collection_name=collection_name,
wait=wait,
keys=keys,
@@ -1137,7 +1160,7 @@ class QdrantClient:
return self.openapi_client.points_api.delete_payload(
collection_name=collection_name,
wait=wait,
delete_payload=rest.DeletePayload(
delete_payload=rest_models.DeletePayload(
keys=keys,
points=_points,
filter=_filter,
@@ -1170,7 +1193,7 @@ class QdrantClient:
points_selector = self._try_argument_to_grpc_selector(points_selector)
return GrpcToRest.convert_update_result(
self._grpc_points_client.ClearPayload(grpc.ClearPayloadPoints(
self.grpc_points.ClearPayload(grpc.ClearPayloadPoints(
collection_name=collection_name,
wait=wait,
points=points_selector
@@ -1210,7 +1233,7 @@ class QdrantClient:
return self.http.collections_api.update_aliases(
timeout=timeout,
change_aliases_operation=rest.ChangeAliasesOperation(
change_aliases_operation=rest_models.ChangeAliasesOperation(
actions=change_aliases_operation
)
)
@@ -1222,8 +1245,8 @@ class QdrantClient:
List of the collections
"""
if self._prefer_grpc:
response = self._grpc_collections_client.List(grpc.ListCollectionsRequest(),
timeout=self._timeout).collections
response = self.grpc_collections.List(grpc.ListCollectionsRequest(),
timeout=self._timeout).collections
return types.CollectionsResponse(
collections=[GrpcToRest.convert_collection_description(description) for description in response]
)
@@ -1241,7 +1264,7 @@ class QdrantClient:
"""
if self._prefer_grpc:
return GrpcToRest.convert_collection_info(
self._grpc_collections_client.Get(grpc.GetCollectionInfoRequest(
self.grpc_collections.Get(grpc.GetCollectionInfoRequest(
collection_name=collection_name
), timeout=self._timeout).result)
return self.http.collections_api.get_collection(collection_name=collection_name).result
@@ -1274,7 +1297,7 @@ class QdrantClient:
return self.http.collections_api.update_collection(
collection_name,
update_collection=rest.UpdateCollection(
update_collection=rest_models.UpdateCollection(
optimizers_config=optimizer_config,
params=collection_params
),
@@ -1362,7 +1385,7 @@ class QdrantClient:
self.delete_collection(collection_name)
create_collection_request = rest.CreateCollection(
create_collection_request = rest_models.CreateCollection(
vectors=vectors_config,
shard_number=shard_number,
replication_factor=replication_factor,
@@ -1469,8 +1492,8 @@ class QdrantClient:
def create_payload_index(self,
collection_name: str,
field_name: str,
field_schema: types.PayloadSchemaType = None,
field_type: types.PayloadSchemaType = None,
field_schema: Optional[types.PayloadSchemaType] = None,
field_type: Optional[types.PayloadSchemaType] = None,
wait: bool = True,
):
"""Creates index for a given payload field.
@@ -1498,7 +1521,7 @@ class QdrantClient:
return self.openapi_client.collections_api.create_field_index(
collection_name=collection_name,
create_field_index=rest.CreateFieldIndex(field_name=field_name, field_schema=field_schema),
create_field_index=rest_models.CreateFieldIndex(field_name=field_name, field_schema=field_schema),
wait=wait
)
@@ -1531,9 +1554,13 @@ class QdrantClient:
Returns:
List of snapshots
"""
return self.openapi_client.collections_api.list_snapshots(collection_name=collection_name).result
snapshots = self.openapi_client.collections_api.list_snapshots(
collection_name=collection_name
).result
assert snapshots is not None, "List snapshots API returned None result"
return snapshots
def create_snapshot(self, collection_name: str) -> types.SnapshotDescription:
def create_snapshot(self, collection_name: str) -> Optional[types.SnapshotDescription]:
"""Create snapshot for a given collection.
Args:
@@ -1552,7 +1579,9 @@ class QdrantClient:
Returns:
List of snapshots
"""
return self.openapi_client.snapshots_api.list_full_snapshots().result
snapshots = self.openapi_client.snapshots_api.list_full_snapshots().result
assert snapshots is not None, "List full snapshots API returned None result"
return snapshots
def create_full_snapshot(self) -> types.SnapshotDescription:
"""Create snapshot for a whole storage.
@@ -1560,10 +1589,16 @@ class QdrantClient:
Returns:
Snapshot description
"""
return self.openapi_client.snapshots_api.create_full_snapshot().result
snapshot_description = self.openapi_client.snapshots_api.create_full_snapshot().result
assert snapshot_description is not None, "Create full snapshot API returned None result"
return snapshot_description
def recover_snapshot(self, collection_name: str, location: str,
priority: Optional[types.SnapshotPriority] = None) -> bool:
def recover_snapshot(
self,
collection_name: str,
location: str,
priority: Optional[types.SnapshotPriority] = None
) -> bool:
"""Recover collection from snapshot.
Args:
@@ -1580,22 +1615,24 @@ class QdrantClient:
Default: `replica`
"""
return self.openapi_client.snapshots_api.recover_from_snapshot(
success = self.openapi_client.snapshots_api.recover_from_snapshot(
collection_name=collection_name,
snapshot_recover=rest.SnapshotRecover(location=location, priority=priority)
).result
assert success is not None, "Recover from snapshot API returned None result"
return success
def lock_storage(self, reason: str):
"""Lock storage for writing.
"""
return self.openapi_client.service_api.post_locks(rest.LocksOption(error_message=reason, write=True))
return self.openapi_client.service_api.post_locks(rest_models.LocksOption(error_message=reason, write=True))
def unlock_storage(self):
"""Unlock storage for writing.
"""
return self.openapi_client.service_api.post_locks(rest.LocksOption(write=False))
return self.openapi_client.service_api.post_locks(rest_models.LocksOption(write=False))
def get_locks(self) -> types.LocksOption:
def get_locks(self) -> Optional[types.LocksOption]:
"""Get current locks state.
"""
return self.openapi_client.service_api.get_locks().result

View File

@@ -40,7 +40,7 @@ class RestBatchUploader(BaseUploader):
def __init__(self, uri, collection_name, **kwargs: Any):
self.collection_name = collection_name
self.openapi_client = SyncApis(host=uri, **kwargs)
self.openapi_client: SyncApis = SyncApis(host=uri, **kwargs)
@classmethod
def start(cls, collection_name=None, uri="http://localhost:6333", **kwargs) -> 'RestBatchUploader':

View File

@@ -2,7 +2,7 @@ import itertools
import math
from abc import ABC
from itertools import islice, count
from typing import Optional, Iterable, Any, Callable, Union, List
from typing import Optional, Iterable, Union, List, Generator
import numpy as np
@@ -25,12 +25,10 @@ def iter_batch(iterable, size) -> Iterable:
class BaseUploader(Worker, ABC):
@classmethod
def iterate_records_batches(cls,
records: Iterable[Record],
batch_size: int
) -> Iterable:
def iterate_records_batches(
cls, records: Iterable[Record], batch_size: int
) -> Iterable:
record_batches = iter_batch(records, batch_size)
for record_batch in record_batches:
@@ -40,25 +38,30 @@ class BaseUploader(Worker, ABC):
yield ids_batch, vectors_batch, payload_batch
@classmethod
def iterate_batches(cls,
vectors: Union[np.ndarray, Iterable[List[float]]],
payload: Optional[Iterable[dict]],
ids: Optional[Iterable[ExtendedPointId]],
batch_size: int,
) -> Iterable:
def iterate_batches(
cls,
vectors: Union[np.ndarray, Iterable[List[float]]],
payload: Optional[Iterable[dict]],
ids: Optional[Iterable[ExtendedPointId]],
batch_size: int,
) -> Iterable:
if ids is None:
ids = itertools.count()
ids_batches = iter_batch(ids, batch_size)
if payload is None:
payload_batches = (None for _ in count())
payload_batches: Union[Generator, Iterable] = (None for _ in count())
else:
payload_batches = iter_batch(payload, batch_size)
if isinstance(vectors, np.ndarray):
num_vectors = vectors.shape[0]
num_batches = int(math.ceil(num_vectors / batch_size))
vector_batches = (vectors[i * batch_size:(i + 1) * batch_size].tolist() for i in range(num_batches))
vector_batches: Union[Generator, Iterable] = (
vectors[i * batch_size : (i + 1) * batch_size].tolist()
for i in range(num_batches)
)
else:
vector_batches = iter_batch(vectors, batch_size)