Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 29 additions & 48 deletions qdrant_client/conversions/conversion.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
import math
import uuid
from datetime import date, datetime, timezone
from typing import Any, Mapping, Sequence, get_args

from google.protobuf.internal.containers import MessageMap
from google.protobuf.json_format import MessageToDict
from google.protobuf.timestamp_pb2 import Timestamp

try:
Expand All @@ -18,22 +18,6 @@
from qdrant_client.conversions.common_types import get_args_subscribed


def has_field(message: Any, field: str) -> bool:
"""
Same as protobuf HasField, but also works for primitive values
(https://stackoverflow.com/questions/51918871/check-if-a-field-has-been-set-in-protocol-buffer-3)

Args:
message (Any): protobuf message
field (str): name of the field
"""
try:
return message.HasField(field)
except ValueError:
all_fields = set([descriptor.name for descriptor, _value in message.ListFields()])
return field in all_fields


# protobuf resolves enum attributes in python, it takes ~0.5us on every access
_NULL_VALUE = NullValue.NULL_VALUE

Expand Down Expand Up @@ -99,43 +83,39 @@ def _json_to_value_kwargs(payload: Any) -> dict[str, Any]:


def value_to_json(value: Value) -> Any:
if isinstance(value, Value):
value_ = MessageToDict(value, preserving_proto_field_name=False)
else:
value_ = value

if "integerValue" in value_:
# by default int are represented as string for precision
# But in python it is OK to just use `int`
return int(value_["integerValue"])
if "doubleValue" in value_:
return value_["doubleValue"]
if "stringValue" in value_:
return value_["stringValue"]
if "boolValue" in value_:
return value_["boolValue"]
if "structValue" in value_:
if "fields" not in value_["structValue"]:
return {}
return dict(
(key, value_to_json(val)) for key, val in value_["structValue"]["fields"].items()
)
if "listValue" in value_:
if "values" in value_["listValue"]:
return list(value_to_json(val) for val in value_["listValue"]["values"])
else:
return []
if "nullValue" in value_:
# Reading the set field directly is several times faster than MessageToDict, which walks the
# message in python
kind = value.WhichOneof("kind")
if kind == "string_value":
return value.string_value
if kind == "integer_value":
return value.integer_value
if kind == "double_value":
double_value = value.double_value
if math.isfinite(double_value):
return double_value
# keep the strings MessageToDict used to return for non-finite doubles
if math.isnan(double_value):
return "NaN"
return "Infinity" if double_value > 0 else "-Infinity"
if kind == "bool_value":
return value.bool_value
if kind == "list_value":
return [value_to_json(val) for val in value.list_value.values]
if kind == "struct_value":
return grpc_to_payload(value.struct_value.fields)
if kind == "null_value":
return None
raise ValueError(f"Not supported value: {value_}") # pragma: no cover
raise ValueError(f"Not supported value: {value}") # pragma: no cover


def payload_to_grpc(payload: dict[str, Any]) -> dict[str, Value]:
return dict((key, json_to_value(val)) for key, val in payload.items())


def grpc_to_payload(grpc_: MessageMap[str, Value]) -> dict[str, Any]:
return dict((key, value_to_json(val)) for key, val in grpc_.items())
# upb implements items() of a map in python, looking up each key is faster
return {key: value_to_json(grpc_[key]) for key in grpc_}


def grpc_payload_schema_to_field_type(model: grpc.PayloadSchemaType) -> grpc.FieldType:
Expand Down Expand Up @@ -671,7 +651,8 @@ def convert_scored_point(cls, model: grpc.ScoredPoint) -> rest.ScoredPoint:
return construct(
rest.ScoredPoint,
id=cls.convert_point_id(model.id),
payload=cls.convert_payload(model.payload) if has_field(model, "payload") else None,
# HasField raises on map fields
payload=cls.convert_payload(model.payload) if len(model.payload) > 0 else None,
score=model.score,
vector=(
cls.convert_vectors_output(model.vectors) if model.HasField("vectors") else None
Expand All @@ -689,7 +670,7 @@ def convert_scored_point(cls, model: grpc.ScoredPoint) -> rest.ScoredPoint:

@classmethod
def convert_payload(cls, model: "MessageMapContainer") -> rest.Payload:
return dict((key, value_to_json(model[key])) for key in model)
return grpc_to_payload(model)

@classmethod
def convert_values_count(cls, model: grpc.ValuesCount) -> rest.ValuesCount:
Expand Down
49 changes: 49 additions & 0 deletions tests/conversions/test_validate_conversions.py
Original file line number Diff line number Diff line change
Expand Up @@ -369,6 +369,55 @@ def test_json_to_value_struct_fields_as_messages(monkeypatch):
assert [conversion.payload_to_grpc(payload) for payload in payloads] == expected


@pytest.mark.parametrize(
"field, value, expected",
[
("integer_value", 2**63 - 1, 2**63 - 1),
("integer_value", -(2**63), -(2**63)),
("double_value", 1.0, 1.0),
("double_value", -0.0, -0.0),
# non-finite doubles are returned as strings, the way MessageToDict renders them
("double_value", float("nan"), "NaN"),
("double_value", float("inf"), "Infinity"),
("double_value", float("-inf"), "-Infinity"),
("string_value", "", ""),
("bool_value", False, False),
("null_value", 0, None),
],
)
def test_value_to_json_scalars(field, value, expected):
from qdrant_client.conversions.conversion import value_to_json
from qdrant_client.grpc import Value

result = value_to_json(Value(**{field: value}))
assert type(result) is type(expected)
assert repr(result) == repr(expected) # tells -0.0 from 0.0


def test_value_to_json_containers():
from qdrant_client.conversions.conversion import json_to_value, value_to_json
from qdrant_client.grpc import ListValue, Struct, Value

assert value_to_json(Value(struct_value=Struct())) == {}
assert value_to_json(Value(list_value=ListValue())) == []

payload = {"a": [1, 1.5, None, {"b": [True, "c", [], {}]}], "d": {"e": {"f": -1}}}
assert value_to_json(json_to_value(payload)) == payload


def test_convert_points_empty_payload():
from qdrant_client import grpc
from qdrant_client.conversions.conversion import GrpcToRest, json_to_value

point_id = grpc.PointId(num=1)
assert GrpcToRest.convert_scored_point(grpc.ScoredPoint(id=point_id)).payload is None
assert GrpcToRest.convert_retrieved_point(grpc.RetrievedPoint(id=point_id)).payload == {}

payload = {"a": json_to_value(1)}
scored_point = GrpcToRest.convert_scored_point(grpc.ScoredPoint(id=point_id, payload=payload))
assert scored_point.payload == {"a": 1}


def test_convert_context_input_flat_pair():
from qdrant_client import models
from qdrant_client.conversions.conversion import GrpcToRest, RestToGrpc
Expand Down
Loading