mirror of
https://github.com/qdrant/qdrant-client.git
synced 2026-07-23 11:11:01 -05:00
new: add py.typed, validate code with mypy (#91)
This commit is contained in:
4
mypy.ini
Normal file
4
mypy.ini
Normal file
@@ -0,0 +1,4 @@
|
||||
[mypy]
|
||||
ignore_missing_imports = True
|
||||
follow_imports = skip
|
||||
exclude = qdrant_client/grpc|qdrant_client/http|tests
|
||||
@@ -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]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
1
qdrant_client/py.typed
Normal file
@@ -0,0 +1 @@
|
||||
partial
|
||||
@@ -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
|
||||
|
||||
@@ -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':
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user