diff --git a/qdrant_client/conversions/conversion.py b/qdrant_client/conversions/conversion.py index 46e03d4a9..3f4d1bf9f 100644 --- a/qdrant_client/conversions/conversion.py +++ b/qdrant_client/conversions/conversion.py @@ -12,7 +12,7 @@ pass from qdrant_client import grpc -from qdrant_client.grpc import ListValue, NullValue, Struct, Value +from qdrant_client.grpc import NullValue, Struct, Value from qdrant_client.http.models import models as rest from qdrant_client._pydantic_compat import construct, to_jsonable_python from qdrant_client.conversions.common_types import get_args_subscribed @@ -34,26 +34,68 @@ def has_field(message: Any, field: str) -> bool: return field in all_fields +# protobuf resolves enum attributes in python, it takes ~0.5us on every access +_NULL_VALUE = NullValue.NULL_VALUE + + def json_to_value(payload: Any) -> Value: - if payload is None: - return Value(null_value=NullValue.NULL_VALUE) - if isinstance(payload, bool): + # exact type checks are cheaper than isinstance and cover almost every value + payload_type = type(payload) + if payload_type is str: + return Value(string_value=payload) + if payload_type is int: + return Value(integer_value=payload) + if payload_type is float: + return Value(double_value=payload) + if payload_type is bool: return Value(bool_value=payload) + if payload is None: + return Value(null_value=_NULL_VALUE) + return Value(**_json_to_value_kwargs(payload)) + + +def _json_to_value_kwargs(payload: Any) -> dict[str, Any]: + # Returns the arguments of the Value constructor instead of a Value. Protobuf builds the whole + # tree from nested dicts of arguments at once, which is much faster than building a Value for + # every element and copying it into its parent. + payload_type = type(payload) + if payload_type is str: + return {"string_value": payload} + if payload_type is int: + return {"integer_value": payload} + if payload_type is float: + return {"double_value": payload} + if payload_type is bool: + return {"bool_value": payload} + if payload is None: + return {"null_value": _NULL_VALUE} + if isinstance(payload, (list, tuple)): + return {"list_value": {"values": [_json_to_value_kwargs(v) for v in payload]}} + if isinstance(payload, dict) and all(isinstance(key, str) for key in payload): + fields = {k: _json_to_struct_field(v) for k, v in payload.items()} + return {"struct_value": {"fields": fields}} + # subclasses, e.g. enums if isinstance(payload, int): - return Value(integer_value=payload) + return {"integer_value": payload} if isinstance(payload, float): - return Value(double_value=payload) + return {"double_value": payload} if isinstance(payload, str): - return Value(string_value=payload) - if isinstance(payload, (list, tuple)): - return Value(list_value=ListValue(values=[json_to_value(v) for v in payload])) - if isinstance(payload, dict): - return Value( - struct_value=Struct(fields=dict((k, json_to_value(v)) for k, v in payload.items())) - ) - if isinstance(payload, datetime) or isinstance(payload, date): - return Value(string_value=to_jsonable_python(payload)) - raise ValueError(f"Not supported json value: {payload}") # pragma: no cover + return {"string_value": payload} + # Encode the rest (uuid, datetime, Decimal, set, bytes, non-str dict keys, etc.) the same way + # the REST client and local mode do + try: + jsonable_payload = to_jsonable_python(payload) + except (KeyError, ValueError) as e: # pydantic v1 raises KeyError, v2 raises ValueError + raise ValueError(f"Not supported json value: {payload}") from e + return _json_to_value_kwargs(jsonable_payload) + + +try: + # protobuf < 6.30 does not accept dicts of arguments as values of a map, e.g. Struct.fields + Struct(fields={"": {}}) + _json_to_struct_field = _json_to_value_kwargs +except TypeError: # pragma: no cover + _json_to_struct_field = json_to_value def value_to_json(value: Value) -> Any: diff --git a/tests/conversions/test_validate_conversions.py b/tests/conversions/test_validate_conversions.py index d7c791ab6..55bad3fde 100644 --- a/tests/conversions/test_validate_conversions.py +++ b/tests/conversions/test_validate_conversions.py @@ -1,9 +1,16 @@ import inspect +import json import logging import re -from datetime import date, datetime, timedelta, timezone +import uuid +from collections import OrderedDict, defaultdict, namedtuple +from datetime import date, datetime, time, timedelta, timezone +from decimal import Decimal +from enum import Enum, IntEnum from inspect import getmembers +from pathlib import Path +import numpy as np import pytest from google.protobuf.json_format import MessageToDict @@ -281,6 +288,87 @@ def test_datetime_to_timestamp_conversions(dt: datetime | date): ), f"Failed for {dt}, should be equal to {grpc_to_rest}" +@pytest.mark.parametrize( + "value", + [ + uuid.UUID("5a6d1c3e-8f0b-4c55-9d6e-0a1b2c3d4e5f"), + Decimal("1.5"), + {1, 2}, + frozenset({"a"}), + b"bytes", + datetime(2021, 1, 1, 12, 30, tzinfo=timezone.utc), + date(2021, 1, 1), + time(12, 30), + timedelta(hours=1), + Path("/tmp/file"), + {1: "non-str key"}, + {"nested": [uuid.UUID(int=1), {"deeper": (Decimal("2"), b"x")}]}, + ], + ids=lambda value: type(value).__name__, +) +def test_json_to_value_matches_rest_encoding(value): + from qdrant_client import models + from qdrant_client.conversions.conversion import json_to_value, value_to_json + from qdrant_client.http.api.points_api import jsonable_encoder + + # the body the REST client sends for the same payload + rest_body = jsonable_encoder(models.SetPayload(payload={"value": value}, points=[1])) + rest_value = json.loads(rest_body)["payload"]["value"] + + assert value_to_json(json_to_value(value)) == rest_value + + +@pytest.mark.parametrize( + "value", [object(), np.float32(1.0), {object(): "key"}], ids=lambda value: type(value).__name__ +) +def test_json_to_value_unsupported(value): + from qdrant_client.conversions.conversion import json_to_value + + with pytest.raises(ValueError, match="Not supported json value"): + json_to_value(value) + + +class IntColor(IntEnum): + RED = 1 + + +class StrColor(str, Enum): + RED = "red" + + +@pytest.mark.parametrize( + "value, plain", + [ + (IntColor.RED, 1), + (StrColor.RED, "red"), + (np.float64(1.5), 1.5), + ({StrColor.RED: [IntColor.RED]}, {"red": [1]}), + (namedtuple("Pair", "x y")(1, "a"), [1, "a"]), + (OrderedDict(a={"b": None}), {"a": {"b": None}}), + (defaultdict(list, a=True), {"a": True}), + ], + ids=lambda value: type(value).__name__, +) +def test_json_to_value_subclasses(value, plain): + from qdrant_client.conversions.conversion import json_to_value + + # subclasses miss the exact type checks, but have to be encoded like the plain values + assert json_to_value(value) == json_to_value(plain) + assert json_to_value({"nested": value}) == json_to_value({"nested": plain}) + + +def test_json_to_value_struct_fields_as_messages(monkeypatch): + from qdrant_client.conversions import conversion + from tests.fixtures.payload import one_random_payload_please + + payloads = [one_random_payload_please(i) for i in range(20)] + expected = [conversion.payload_to_grpc(payload) for payload in payloads] + + # the path taken with protobuf < 6.30, which does not accept dicts as values of a map + monkeypatch.setattr(conversion, "_json_to_struct_field", conversion.json_to_value) + assert [conversion.payload_to_grpc(payload) for payload in payloads] == expected + + def test_convert_context_input_flat_pair(): from qdrant_client import models from qdrant_client.conversions.conversion import GrpcToRest, RestToGrpc