mirror of
https://github.com/qdrant/qdrant-client.git
synced 2026-09-21 13:37:55 -05:00
fix: remove enter/exit, add conversion tests
This commit is contained in:
@@ -97,6 +97,7 @@ class AsyncQdrantServerless:
|
||||
if self._grpc_collections is None:
|
||||
self._remote._init_grpc_channel()
|
||||
self._grpc_collections = CollectionsServiceStub(self._remote._grpc_channel_pool[0])
|
||||
assert self._grpc_collections is not None
|
||||
return self._grpc_collections
|
||||
|
||||
def _collections_timeout(self, timeout: Optional[int]) -> int:
|
||||
|
||||
@@ -89,6 +89,7 @@ class QdrantServerless:
|
||||
# reuse the delegate's channel: same host, tls, api-key metadata and options
|
||||
self._remote._init_grpc_channel()
|
||||
self._grpc_collections = CollectionsServiceStub(self._remote._grpc_channel_pool[0])
|
||||
assert self._grpc_collections is not None
|
||||
return self._grpc_collections
|
||||
|
||||
def _collections_timeout(self, timeout: Optional[int]) -> int:
|
||||
@@ -106,12 +107,6 @@ class QdrantServerless:
|
||||
self._grpc_collections = None
|
||||
self._remote.close(grpc_grace=grpc_grace, **kwargs)
|
||||
|
||||
def __enter__(self) -> "QdrantServerless":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: Any) -> None:
|
||||
self.close()
|
||||
|
||||
# region collections
|
||||
|
||||
def create_collection(
|
||||
|
||||
@@ -52,7 +52,7 @@ def dense_vector_to_grpc(model: models.DenseVectorConfig) -> pb2.DenseVectorConf
|
||||
if precision_tier is not None:
|
||||
result.precision_tier = _PRECISION_TO_GRPC[precision_tier]
|
||||
return result
|
||||
case _:
|
||||
case _: # pragma: no cover
|
||||
raise ValueError(f"Unexpected DenseVectorConfig shape: {model!r}")
|
||||
|
||||
|
||||
@@ -81,7 +81,7 @@ def sparse_vector_to_grpc(model: models.SparseVectorConfig) -> pb2.SparseVectorC
|
||||
if precision_tier is not None:
|
||||
result.precision_tier = _PRECISION_TO_GRPC[precision_tier]
|
||||
return result
|
||||
case _:
|
||||
case _: # pragma: no cover
|
||||
raise ValueError(f"Unexpected SparseVectorConfig shape: {model!r}")
|
||||
|
||||
|
||||
@@ -99,7 +99,7 @@ def _stopwords_to_grpc(model: models.StopwordsSet) -> pb2.StopwordsSet:
|
||||
match model:
|
||||
case models.StopwordsSet(languages=languages, custom=custom):
|
||||
return pb2.StopwordsSet(languages=list(languages), custom=list(custom))
|
||||
case _:
|
||||
case _: # pragma: no cover
|
||||
raise ValueError(f"Unexpected StopwordsSet shape: {model!r}")
|
||||
|
||||
|
||||
@@ -118,14 +118,14 @@ def _stemmer_to_grpc(model: models.StemmingAlgorithm) -> pb2.StemmingAlgorithm:
|
||||
match snowball:
|
||||
case models.SnowballParams(language=language):
|
||||
result.snowball.language = language
|
||||
case _:
|
||||
case _: # pragma: no cover
|
||||
raise ValueError(f"Unexpected SnowballParams shape: {snowball!r}")
|
||||
elif disabled:
|
||||
result.disabled.SetInParent()
|
||||
else:
|
||||
raise ValueError("StemmingAlgorithm requires either snowball or disabled=True")
|
||||
return result
|
||||
case _:
|
||||
case _: # pragma: no cover
|
||||
raise ValueError(f"Unexpected StemmingAlgorithm shape: {model!r}")
|
||||
|
||||
|
||||
@@ -137,7 +137,7 @@ def _stemmer_from_grpc(grpc_model: pb2.StemmingAlgorithm) -> models.StemmingAlgo
|
||||
)
|
||||
if kind == "disabled":
|
||||
return models.StemmingAlgorithm(disabled=True)
|
||||
raise ValueError(f"Unknown stemming_params variant: {kind}")
|
||||
raise ValueError(f"Unknown stemming_params variant: {kind}") # pragma: no cover
|
||||
|
||||
|
||||
def payload_index_to_grpc(model: models.PayloadIndex) -> pb2.PayloadIndexConfig:
|
||||
@@ -149,7 +149,7 @@ def payload_index_to_grpc(model: models.PayloadIndex) -> pb2.PayloadIndexConfig:
|
||||
match prefix:
|
||||
case models.KeywordPrefixParams():
|
||||
result.keyword.prefix.SetInParent()
|
||||
case _:
|
||||
case _: # pragma: no cover
|
||||
raise ValueError(f"Unexpected KeywordPrefixParams shape: {prefix!r}")
|
||||
case models.IntegerIndex(type=_type, lookup=lookup, range=range_):
|
||||
result.integer.SetInParent()
|
||||
@@ -195,7 +195,7 @@ def payload_index_to_grpc(model: models.PayloadIndex) -> pb2.PayloadIndexConfig:
|
||||
result.geo.SetInParent()
|
||||
case models.BoolIndex(type=_type):
|
||||
result.bool.SetInParent()
|
||||
case _:
|
||||
case _: # pragma: no cover
|
||||
raise ValueError(f"Unknown payload index type: {model}")
|
||||
return result
|
||||
|
||||
@@ -246,7 +246,7 @@ def payload_index_from_grpc(grpc_model: pb2.PayloadIndexConfig) -> models.Payloa
|
||||
if kind == "bool":
|
||||
_bool = grpc_model.bool
|
||||
return models.BoolIndex()
|
||||
raise ValueError(f"Unknown payload index type: {kind}")
|
||||
raise ValueError(f"Unknown payload index type: {kind}") # pragma: no cover
|
||||
|
||||
|
||||
def collection_config_to_grpc(model: models.CollectionConfig) -> pb2.CollectionConfig:
|
||||
@@ -264,7 +264,7 @@ def collection_config_to_grpc(model: models.CollectionConfig) -> pb2.CollectionC
|
||||
for field, index in payload_indexes.items():
|
||||
result.payload_indexes[field].CopyFrom(payload_index_to_grpc(index))
|
||||
return result
|
||||
case _:
|
||||
case _: # pragma: no cover
|
||||
raise ValueError(f"Unexpected CollectionConfig shape: {model!r}")
|
||||
|
||||
|
||||
|
||||
@@ -15,31 +15,6 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from qdrant_client.http.models import Distance, TokenizerType
|
||||
|
||||
__all__ = [
|
||||
"Distance",
|
||||
"TokenizerType",
|
||||
"PrecisionTier",
|
||||
"DenseVectorConfig",
|
||||
"SparseVectorConfig",
|
||||
"KeywordPrefixParams",
|
||||
"KeywordIndex",
|
||||
"IntegerIndex",
|
||||
"FloatIndex",
|
||||
"UuidIndex",
|
||||
"DatetimeIndex",
|
||||
"StopwordsSet",
|
||||
"SnowballParams",
|
||||
"StemmingAlgorithm",
|
||||
"TextIndex",
|
||||
"GeoIndex",
|
||||
"BoolIndex",
|
||||
"PayloadIndex",
|
||||
"CollectionConfig",
|
||||
"CollectionInfo",
|
||||
"CollectionSummary",
|
||||
"CollectionsList",
|
||||
]
|
||||
|
||||
|
||||
class PrecisionTier(str, Enum):
|
||||
"""How much vector precision may be traded for cost.
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
from google.protobuf.message import Message
|
||||
|
||||
from qdrant_client.serverless.grpc import serverless_collections_pb2 as pb2
|
||||
|
||||
# Dense vectors: every Distance and every PrecisionTier member is covered, so a
|
||||
# mis-mapped enum constant cannot hide behind a self-consistent round-trip
|
||||
# (_*_FROM_GRPC is derived by inverting _*_TO_GRPC).
|
||||
dense_vector = pb2.DenseVectorConfig(size=1536, distance=pb2.COSINE)
|
||||
dense_vector_multivector = pb2.DenseVectorConfig(
|
||||
size=128, distance=pb2.DOT, multivector=True, precision_tier=pb2.LOW
|
||||
)
|
||||
dense_vector_manhattan = pb2.DenseVectorConfig(
|
||||
size=4, distance=pb2.MANHATTAN, precision_tier=pb2.MEDIUM
|
||||
)
|
||||
dense_vector_euclid = pb2.DenseVectorConfig(size=4, distance=pb2.EUCLID, precision_tier=pb2.HIGH)
|
||||
|
||||
sparse_vector = pb2.SparseVectorConfig(use_idf=True)
|
||||
sparse_vector_tier = pb2.SparseVectorConfig(use_idf=False, precision_tier=pb2.HIGH)
|
||||
|
||||
payload_index_keyword = pb2.PayloadIndexConfig(keyword=pb2.KeywordIndex())
|
||||
payload_index_integer = pb2.PayloadIndexConfig(integer=pb2.IntegerIndex())
|
||||
# `range=False` is set, not absent: catches `if range_:` in place of `is not None`
|
||||
payload_index_integer_falsy = pb2.PayloadIndexConfig(
|
||||
integer=pb2.IntegerIndex(lookup=True, range=False)
|
||||
)
|
||||
payload_index_float = pb2.PayloadIndexConfig(float=pb2.FloatIndex())
|
||||
payload_index_uuid = pb2.PayloadIndexConfig(uuid=pb2.UuidIndex())
|
||||
payload_index_datetime = pb2.PayloadIndexConfig(datetime=pb2.DatetimeIndex())
|
||||
payload_index_text = pb2.PayloadIndexConfig(text=pb2.TextIndex())
|
||||
# every optional TextIndex field set, with falsy values where they are legal
|
||||
payload_index_text_full = pb2.PayloadIndexConfig(
|
||||
text=pb2.TextIndex(
|
||||
tokenizer=pb2.MULTILINGUAL,
|
||||
lowercase=False,
|
||||
phrase_matching=True,
|
||||
min_token_len=0,
|
||||
max_token_len=20,
|
||||
)
|
||||
)
|
||||
payload_index_text_prefix = pb2.PayloadIndexConfig(text=pb2.TextIndex(tokenizer=pb2.PREFIX))
|
||||
payload_index_text_whitespace = pb2.PayloadIndexConfig(
|
||||
text=pb2.TextIndex(tokenizer=pb2.WHITESPACE)
|
||||
)
|
||||
payload_index_text_word = pb2.PayloadIndexConfig(text=pb2.TextIndex(tokenizer=pb2.WORD))
|
||||
payload_index_geo = pb2.PayloadIndexConfig(geo=pb2.GeoIndex())
|
||||
payload_index_bool = pb2.PayloadIndexConfig(bool=pb2.BoolIndex())
|
||||
|
||||
stopwords_empty = pb2.StopwordsSet()
|
||||
stopwords_languages = pb2.StopwordsSet(languages=["english", "german"])
|
||||
stopwords_custom = pb2.StopwordsSet(languages=["english"], custom=["foo", "bar"])
|
||||
|
||||
stemmer_snowball = pb2.StemmingAlgorithm(snowball=pb2.SnowballParams(language="english"))
|
||||
stemmer_disabled = pb2.StemmingAlgorithm(disabled=pb2.DisabledStemmer())
|
||||
|
||||
# presence on an empty submessage: `prefix` set but carrying no fields
|
||||
payload_index_keyword_prefix = pb2.PayloadIndexConfig(
|
||||
keyword=pb2.KeywordIndex(prefix=pb2.KeywordPrefixParams())
|
||||
)
|
||||
payload_index_text_analysis = pb2.PayloadIndexConfig(
|
||||
text=pb2.TextIndex(
|
||||
ascii_folding=True,
|
||||
stopwords=stopwords_custom,
|
||||
stemmer=stemmer_snowball,
|
||||
)
|
||||
)
|
||||
payload_index_text_stemmer_disabled = pb2.PayloadIndexConfig(
|
||||
text=pb2.TextIndex(
|
||||
ascii_folding=False,
|
||||
stopwords=stopwords_empty,
|
||||
stemmer=stemmer_disabled,
|
||||
)
|
||||
)
|
||||
|
||||
collection_config = pb2.CollectionConfig(
|
||||
dense_vectors={"": dense_vector},
|
||||
payload_indexes={"user_id": payload_index_keyword},
|
||||
)
|
||||
collection_config_named_vectors = pb2.CollectionConfig(
|
||||
dense_vectors={"dense": dense_vector, "colbert": dense_vector_multivector},
|
||||
sparse_vectors={"bm25": sparse_vector},
|
||||
payload_indexes={"age": payload_index_integer_falsy, "text": payload_index_text_full},
|
||||
)
|
||||
collection_config_empty = pb2.CollectionConfig()
|
||||
|
||||
fixtures: dict[str, list[Message]] = {
|
||||
"dense_vector": [
|
||||
dense_vector,
|
||||
dense_vector_multivector,
|
||||
dense_vector_manhattan,
|
||||
dense_vector_euclid,
|
||||
],
|
||||
"sparse_vector": [
|
||||
sparse_vector,
|
||||
sparse_vector_tier,
|
||||
],
|
||||
"stopwords": [
|
||||
stopwords_empty,
|
||||
stopwords_languages,
|
||||
stopwords_custom,
|
||||
],
|
||||
"stemmer": [
|
||||
stemmer_snowball,
|
||||
stemmer_disabled,
|
||||
],
|
||||
"payload_index": [
|
||||
payload_index_keyword,
|
||||
payload_index_keyword_prefix,
|
||||
payload_index_integer,
|
||||
payload_index_integer_falsy,
|
||||
payload_index_float,
|
||||
payload_index_uuid,
|
||||
payload_index_datetime,
|
||||
payload_index_text,
|
||||
payload_index_text_full,
|
||||
payload_index_text_prefix,
|
||||
payload_index_text_whitespace,
|
||||
payload_index_text_word,
|
||||
payload_index_text_analysis,
|
||||
payload_index_text_stemmer_disabled,
|
||||
payload_index_geo,
|
||||
payload_index_bool,
|
||||
],
|
||||
"collection_config": [
|
||||
collection_config,
|
||||
collection_config_named_vectors,
|
||||
collection_config_empty,
|
||||
],
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
from inspect import getmembers, isfunction
|
||||
from typing import Any, Callable, get_args
|
||||
|
||||
import pytest
|
||||
from google.protobuf.json_format import MessageToDict
|
||||
|
||||
from qdrant_client._pydantic_compat import model_fields
|
||||
from qdrant_client.serverless import conversions, models
|
||||
from qdrant_client.serverless.conversions import (
|
||||
_DISTANCE_TO_GRPC,
|
||||
_PRECISION_TO_GRPC,
|
||||
_TOKENIZER_TO_GRPC,
|
||||
)
|
||||
from qdrant_client.serverless.grpc import serverless_collections_pb2 as pb2
|
||||
from tests.conversions.serverless_fixtures import (
|
||||
fixtures as class_fixtures,
|
||||
)
|
||||
|
||||
|
||||
def _converters(suffix: str) -> dict[str, Callable[[Any], Any]]:
|
||||
"""Map converters by the model they convert, keyed independently of privacy.
|
||||
|
||||
Sub-message converters are module-private (`_stemmer_to_grpc`), so the
|
||||
leading underscore is stripped: a fixture key stays valid if a converter
|
||||
later becomes public, or the reverse.
|
||||
"""
|
||||
matched = [
|
||||
(name, func) for name, func in getmembers(conversions, isfunction) if name.endswith(suffix)
|
||||
]
|
||||
converters = {name[: -len(suffix)].lstrip("_"): func for name, func in matched}
|
||||
assert len(converters) == len(matched), f"stem collision among {suffix} converters"
|
||||
return converters
|
||||
|
||||
|
||||
def test_conversion_completeness() -> None:
|
||||
"""Round-trip every fixture grpc -> model -> grpc, as tests/conversions does.
|
||||
|
||||
Starting from the grpc side is what exercises server-authored messages; a
|
||||
model -> grpc -> model round-trip cannot detect a mis-mapped enum, because
|
||||
the decode maps are derived by inverting the encode maps.
|
||||
"""
|
||||
to_grpc, from_grpc = _converters("_to_grpc"), _converters("_from_grpc")
|
||||
|
||||
assert set(to_grpc) == set(from_grpc), "every converter needs both directions"
|
||||
assert set(to_grpc) == set(class_fixtures), "every converter needs a fixture"
|
||||
|
||||
for model_name, fixtures in class_fixtures.items():
|
||||
for fixture in fixtures:
|
||||
model_fixture = from_grpc[model_name](fixture)
|
||||
|
||||
back_convert_function_name = f"{model_name}_to_grpc"
|
||||
|
||||
print(
|
||||
f"back_convert_function_name: {back_convert_function_name} for {type(model_fixture)}"
|
||||
)
|
||||
|
||||
grpc_fixture = to_grpc[model_name](model_fixture)
|
||||
assert MessageToDict(grpc_fixture) == MessageToDict(
|
||||
fixture
|
||||
), f"{model_name} conversion is broken for {fixture}"
|
||||
|
||||
|
||||
def test_every_payload_index_kind_has_a_fixture() -> None:
|
||||
"""Fails when the proto gains an index kind that nothing covers."""
|
||||
oneof = pb2.PayloadIndexConfig.DESCRIPTOR.oneofs_by_name["index"]
|
||||
covered = {fixture.WhichOneof("index") for fixture in class_fixtures["payload_index"]}
|
||||
assert covered == {field.name for field in oneof.fields}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"mapping,descriptor",
|
||||
[
|
||||
(_DISTANCE_TO_GRPC, pb2.Distance.DESCRIPTOR),
|
||||
(_PRECISION_TO_GRPC, pb2.PrecisionTier.DESCRIPTOR),
|
||||
(_TOKENIZER_TO_GRPC, pb2.Tokenizer.DESCRIPTOR),
|
||||
],
|
||||
)
|
||||
def test_enum_maps_match_proto(mapping: dict, descriptor: Any) -> None:
|
||||
"""Pin each enum mapping by name, and require every proto member to be mapped.
|
||||
|
||||
The round-trip alone cannot catch a swapped pair of constants; comparing the
|
||||
proto member name against the model member name can.
|
||||
"""
|
||||
proto_names = {
|
||||
value.name for value in descriptor.values if not value.name.endswith("_UNSPECIFIED")
|
||||
}
|
||||
assert proto_names == {descriptor.values_by_number[value].name for value in mapping.values()}
|
||||
|
||||
for model_value, grpc_value in mapping.items():
|
||||
assert descriptor.values_by_number[grpc_value].name == model_value.name
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,message",
|
||||
[
|
||||
(models.DenseVectorConfig, pb2.DenseVectorConfig),
|
||||
(models.SparseVectorConfig, pb2.SparseVectorConfig),
|
||||
(models.CollectionConfig, pb2.CollectionConfig),
|
||||
(models.KeywordIndex, pb2.KeywordIndex),
|
||||
(models.KeywordPrefixParams, pb2.KeywordPrefixParams),
|
||||
(models.IntegerIndex, pb2.IntegerIndex),
|
||||
(models.TextIndex, pb2.TextIndex),
|
||||
(models.StopwordsSet, pb2.StopwordsSet),
|
||||
(models.SnowballParams, pb2.SnowballParams),
|
||||
(models.StemmingAlgorithm, pb2.StemmingAlgorithm),
|
||||
# response shapes: a dropped field here silently loses server data
|
||||
(models.CollectionInfo, pb2.GetCollectionResponse),
|
||||
(models.CollectionSummary, pb2.CollectionSummary),
|
||||
(models.CollectionsList, pb2.ListCollectionsResponse),
|
||||
],
|
||||
)
|
||||
def test_models_match_proto_messages(model: Any, message: Any) -> None:
|
||||
"""A field added on either side has to be added on the other."""
|
||||
# `type` is the client-side union tag, it has no proto counterpart
|
||||
assert set(model_fields(model)) - {"type"} == {
|
||||
field.name for field in message.DESCRIPTOR.fields
|
||||
}
|
||||
|
||||
|
||||
def test_payload_index_union_covers_every_model() -> None:
|
||||
kinds = [index_type().type for index_type in get_args(models.PayloadIndex)]
|
||||
assert sorted(kinds) == sorted(
|
||||
field.name for field in pb2.PayloadIndexConfig.DESCRIPTOR.oneofs_by_name["index"].fields
|
||||
)
|
||||
@@ -5,4 +5,7 @@ set -ex
|
||||
coverage run --include='qdrant_client/conversions/conversion.py' -m pytest tests/conversions/test_validate_conversions.py -vv -s
|
||||
coverage report --fail-under=98
|
||||
|
||||
coverage run --include='qdrant_client/serverless/conversions.py' -m pytest tests/conversions/test_validate_serverless_conversions.py -vv -s
|
||||
coverage report --fail-under=98
|
||||
|
||||
#coverage html
|
||||
|
||||
Reference in New Issue
Block a user