fix: remove enter/exit, add conversion tests

This commit is contained in:
George Panchuk
2026-09-15 07:35:03 +02:00
committed by GitHub
parent 75bdd8edcc
commit f9bdad7cf4
7 changed files with 267 additions and 41 deletions
+1
View File
@@ -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:
+1 -6
View File
@@ -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(
+10 -10
View File
@@ -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}")
-25
View File
@@ -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.
+128
View File
@@ -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
)
+3
View File
@@ -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