diff --git a/.librarian/generator-input/client-post-processing/firestore-integration.yaml b/.librarian/generator-input/client-post-processing/firestore-integration.yaml index f43581fb6130..dba36740817e 100644 --- a/.librarian/generator-input/client-post-processing/firestore-integration.yaml +++ b/.librarian/generator-input/client-post-processing/firestore-integration.yaml @@ -78,6 +78,7 @@ replacements: BSONObjectId, BSONRegex, BSONTimestamp, + BSONType, ) from google.cloud.firestore_v1.client import Client from google.cloud.firestore_v1.collection import CollectionReference @@ -188,6 +189,7 @@ replacements: "BSONObjectId", "BSONRegex", "BSONTimestamp", + "BSONType", "Client", "CountAggregation", "CollectionGroup", @@ -268,6 +270,7 @@ replacements: BSONObjectId, BSONRegex, BSONTimestamp, + BSONType, Client, CollectionGroup, CollectionReference, @@ -333,6 +336,7 @@ replacements: "BSONObjectId", "BSONRegex", "BSONTimestamp", + "BSONType", "Client", "CountAggregation", "CollectionGroup", diff --git a/packages/google-cloud-firestore/google/cloud/firestore/__init__.py b/packages/google-cloud-firestore/google/cloud/firestore/__init__.py index f14aa807d509..44a9651d6c02 100644 --- a/packages/google-cloud-firestore/google/cloud/firestore/__init__.py +++ b/packages/google-cloud-firestore/google/cloud/firestore/__init__.py @@ -43,6 +43,7 @@ BSONObjectId, BSONRegex, BSONTimestamp, + BSONType, Client, CollectionGroup, CollectionReference, @@ -108,6 +109,7 @@ "BSONObjectId", "BSONRegex", "BSONTimestamp", + "BSONType", "Client", "CountAggregation", "CollectionGroup", diff --git a/packages/google-cloud-firestore/google/cloud/firestore_v1/__init__.py b/packages/google-cloud-firestore/google/cloud/firestore_v1/__init__.py index f8a91acf9124..4656b5b5441f 100644 --- a/packages/google-cloud-firestore/google/cloud/firestore_v1/__init__.py +++ b/packages/google-cloud-firestore/google/cloud/firestore_v1/__init__.py @@ -55,6 +55,7 @@ BSONObjectId, BSONRegex, BSONTimestamp, + BSONType, ) from google.cloud.firestore_v1.client import Client from google.cloud.firestore_v1.collection import CollectionReference @@ -165,6 +166,7 @@ "BSONObjectId", "BSONRegex", "BSONTimestamp", + "BSONType", "Client", "CountAggregation", "CollectionGroup", diff --git a/packages/google-cloud-firestore/google/cloud/firestore_v1/_helpers.py b/packages/google-cloud-firestore/google/cloud/firestore_v1/_helpers.py index 9793c0685121..afed5b2a703c 100644 --- a/packages/google-cloud-firestore/google/cloud/firestore_v1/_helpers.py +++ b/packages/google-cloud-firestore/google/cloud/firestore_v1/_helpers.py @@ -44,7 +44,7 @@ import google from google.cloud import exceptions # type: ignore from google.cloud.firestore_v1 import transforms, types -from google.cloud.firestore_v1.bson import _BSONType +from google.cloud.firestore_v1.bson import BSONType from google.cloud.firestore_v1.field_path import FieldPath, parse_field_path from google.cloud.firestore_v1.types import common, document, write from google.cloud.firestore_v1.types.write import DocumentTransform @@ -211,7 +211,7 @@ def encode_value(value) -> types.document.Value: if document_path is not None: return document.Value(reference_value=document_path) - if isinstance(value, _BSONType): + if isinstance(value, BSONType): return encode_value(value._to_map_value()) if isinstance(value, GeoPoint): @@ -350,7 +350,18 @@ def reference_value_to_document(reference_value, client) -> Any: def decode_value( value, client ) -> Union[ - None, bool, int, float, list, datetime.datetime, str, bytes, dict, GeoPoint, Vector + None, + bool, + int, + float, + list, + datetime.datetime, + str, + bytes, + dict, + GeoPoint, + Vector, + BSONType, ]: """Converts a Firestore protobuf ``Value`` to a native Python value. @@ -362,7 +373,9 @@ def decode_value( Returns: Union[NoneType, bool, int, float, datetime.datetime, \ - str, bytes, dict, ~google.cloud.Firestore.GeoPoint]: A native + str, bytes, dict, ~google.cloud.Firestore.GeoPoint, \ + ~google.cloud.firestore_v1.vector.Vector, \ + ~google.cloud.firestore_v1.bson.BSONType]: A native \ Python value converted from the ``value``. Raises: @@ -402,7 +415,10 @@ def decode_value( raise ValueError("Unknown ``value_type``", value_type) -def decode_dict(value_fields, client) -> Union[dict, Vector]: +def decode_dict( + value_fields, + client, +) -> Union[dict, Vector, BSONType, bytes]: """Converts a protobuf map of Firestore ``Value``-s. Args: @@ -412,9 +428,9 @@ def decode_dict(value_fields, client) -> Union[dict, Vector]: A client that has a document factory. Returns: - Dict[str, Union[NoneType, bool, int, float, datetime.datetime, \ - str, bytes, dict, ~google.cloud.Firestore.GeoPoint]]: A dictionary - of native Python values converted from the ``value_fields``. + Union[dict, ~google.cloud.firestore_v1.vector.Vector, \ + ~google.cloud.firestore_v1.bson.BSONType, bytes]: A dictionary of native \ + Python values, Vector, BSON object, or bytes converted from ``value_fields``. """ value_fields_pb = getattr(value_fields, "_pb", value_fields) res = {key: decode_value(value, client) for key, value in value_fields_pb.items()} @@ -425,6 +441,10 @@ def decode_dict(value_fields, client) -> Union[dict, Vector]: values = cast(Sequence[float], res["value"]) return Vector(values) + decoded = BSONType._from_dict(res) + if decoded is not None: + return decoded + return res diff --git a/packages/google-cloud-firestore/google/cloud/firestore_v1/bson.py b/packages/google-cloud-firestore/google/cloud/firestore_v1/bson.py index 13bdcac7ffa1..b1a15e91b248 100644 --- a/packages/google-cloud-firestore/google/cloud/firestore_v1/bson.py +++ b/packages/google-cloud-firestore/google/cloud/firestore_v1/bson.py @@ -27,9 +27,10 @@ import abc import decimal import re -from typing import Any, Dict, Union +from typing import Any, Callable, Dict, Optional, Union __all__ = [ + "BSONType", "BSONObjectId", "BSONMinKey", "BSONMaxKey", @@ -44,7 +45,7 @@ _HEX_24_REGEX = re.compile(r"^[0-9a-fA-F]{24}$") -class _BSONType(abc.ABC): +class BSONType(abc.ABC): """Abstract base class for all BSON type containers in Firestore.""" __slots__ = () @@ -65,11 +66,33 @@ def __eq__(self, other: Any) -> bool: def __hash__(self) -> int: """Hash representation contract for set and dictionary keys.""" + @classmethod + def _from_dict(cls, data: Any) -> Optional[Union["BSONType", bytes]]: + """Deserializes a BSON wire map dictionary into a BSON instance or bytes. + + Args: + data (Any): Potential BSON wire map dictionary. + + Returns: + Optional[Union[BSONType, bytes]]: Deserialized BSON container + instance or bytes, or None if not a BSON wire map or if decoding fails. + """ + if not isinstance(data, dict) or len(data) != 1: + return None + key, val = next(iter(data.items())) + decoder = _BSON_DECODERS.get(key) + if decoder is None: + return None + try: + return decoder(val) + except Exception: + return None + def __repr__(self) -> str: return f"{self.__class__.__name__}()" -class BSONObjectId(_BSONType): +class BSONObjectId(BSONType): """Represents a 12-byte BSON ObjectId identifier. Args: @@ -128,7 +151,7 @@ def __hash__(self) -> int: return hash((type(self), self._value)) -class BSONMinKey(_BSONType): +class BSONMinKey(BSONType): """Represents the BSON MinKey sentinel value for query range boundaries.""" __slots__ = () @@ -146,7 +169,7 @@ def __hash__(self) -> int: return hash(type(self)) -class BSONMaxKey(_BSONType): +class BSONMaxKey(BSONType): """Represents the BSON MaxKey sentinel value for query range boundaries.""" __slots__ = () @@ -164,7 +187,7 @@ def __hash__(self) -> int: return hash(type(self)) -class BSONInt32(_BSONType): +class BSONInt32(BSONType): """Represents a 32-bit signed integer value container for Firestore BSON. Args: @@ -221,7 +244,7 @@ def __hash__(self) -> int: return hash((type(self), self._value)) -class BSONBinary(_BSONType): +class BSONBinary(BSONType): """Represents a BSON binary data container with a subtype for Firestore. Args: @@ -286,7 +309,7 @@ def __hash__(self) -> int: return hash((type(self), self._data, self._subtype)) -class BSONTimestamp(_BSONType): +class BSONTimestamp(BSONType): """Container for BSON Timestamp values. Args: @@ -347,7 +370,7 @@ def __hash__(self) -> int: return hash((type(self), self._seconds, self._increment)) -class BSONRegex(_BSONType): +class BSONRegex(BSONType): """Represents a BSON Regular Expression container for Firestore. Args: @@ -409,7 +432,7 @@ def __hash__(self) -> int: return hash((type(self), self._pattern, self._options)) -class BSONDecimal128(_BSONType): +class BSONDecimal128(BSONType): """Represents a BSON 128-bit Decimal container for Firestore. Args: @@ -507,3 +530,21 @@ def __hash__(self) -> int: return hash(d) except decimal.InvalidOperation: return hash((type(self), self._value)) + + +_BSON_DECODERS: Dict[str, Callable[..., Optional[Union[BSONType, bytes]]]] = { + "__oid__": BSONObjectId, + "__min__": lambda _: BSONMinKey(), + "__max__": lambda _: BSONMaxKey(), + "__int__": BSONInt32, + "__decimal128__": BSONDecimal128, + "__binary__": lambda v: (v[1:] if v[0] == 0 else BSONBinary(v[1:], subtype=v[0])) + if isinstance(v, (bytes, bytearray)) and len(v) >= 1 + else None, + "__request_timestamp__": lambda v: BSONTimestamp(v["seconds"], v["increment"]) + if isinstance(v, dict) and "seconds" in v and "increment" in v + else None, + "__regex__": lambda v: BSONRegex(v["pattern"], v.get("options", "")) + if isinstance(v, dict) and "pattern" in v + else None, +} diff --git a/packages/google-cloud-firestore/google/cloud/firestore_v1/order.py b/packages/google-cloud-firestore/google/cloud/firestore_v1/order.py index a3d65cc5000e..dab281d82fcc 100644 --- a/packages/google-cloud-firestore/google/cloud/firestore_v1/order.py +++ b/packages/google-cloud-firestore/google/cloud/firestore_v1/order.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import decimal import math from enum import Enum from typing import Any @@ -19,6 +20,45 @@ from google.cloud.firestore_v1._helpers import GeoPoint, decode_value +def _to_number(val: Any) -> Any: + """Extract a numeric value (int, float, Decimal) from a Value protobuf or Python value. + + Directly inspects the protobuf value_type without calling decode_value() + for optimal performance. + """ + value_pb = getattr(val, "_pb", val) + which = ( + value_pb.WhichOneof("value_type") if hasattr(value_pb, "WhichOneof") else None + ) + + if which == "integer_value": + return value_pb.integer_value + elif which == "double_value": + return value_pb.double_value + elif which == "map_value": + fields = value_pb.map_value.fields + if "__int__" in fields: + return fields["__int__"].integer_value + elif "__decimal128__" in fields: + return decimal.Decimal(fields["__decimal128__"].string_value) + + num = decode_value(val, None) + to_decimal = getattr(num, "to_decimal", None) + return to_decimal() if callable(to_decimal) else getattr(num, "value", num) + + +def _is_nan(val: Any) -> bool: + """Check if a numeric value is NaN, safely handling OverflowError and non-floats.""" + if hasattr(val, "is_nan"): + return val.is_nan() + if isinstance(val, (int, decimal.Decimal)): + return False + try: + return math.isnan(val) + except (TypeError, OverflowError): + return False + + class TypeOrder(Enum): """The supported Data Type. @@ -36,10 +76,17 @@ class TypeOrder(Enum): ARRAY = 8 OBJECT = 9 VECTOR = 10 + BSON_MIN_KEY = 11 + BSON_MAX_KEY = 12 + BSON_OBJECT_ID = 13 + BSON_BINARY = 14 + BSON_REGEX = 15 + BSON_TIMESTAMP = 16 @staticmethod def from_value(value) -> Any: - v = value._pb.WhichOneof("value_type") + value_pb = getattr(value, "_pb", value) + v = value_pb.WhichOneof("value_type") lut = { "null_value": TypeOrder.NULL, "boolean_value": TypeOrder.BOOLEAN, @@ -58,27 +105,50 @@ def from_value(value) -> Any: raise ValueError(f"Could not detect value type for {v}") if v == "map_value": - if ( - "__type__" in value.map_value.fields - and value.map_value.fields["__type__"].string_value == "__vector__" - ): + fields = value_pb.map_value.fields + if len(fields) == 1: + key = next(iter(fields)) + bson_order = _BSON_KEY_TO_TYPE_ORDER.get(key) + if bson_order is not None: + return bson_order + if "__type__" in fields and fields["__type__"].string_value == "__vector__": return TypeOrder.VECTOR return lut[v] +# Maps BSON wire map keys directly to their corresponding TypeOrder. +# BSONInt32 and BSONDecimal128 map to TypeOrder.NUMBER, enabling cross-type comparisons. +_BSON_KEY_TO_TYPE_ORDER = { + "__min__": TypeOrder.BSON_MIN_KEY, + "__max__": TypeOrder.BSON_MAX_KEY, + "__oid__": TypeOrder.BSON_OBJECT_ID, + "__int__": TypeOrder.NUMBER, + "__decimal128__": TypeOrder.NUMBER, + "__binary__": TypeOrder.BSON_BINARY, + "__request_timestamp__": TypeOrder.BSON_TIMESTAMP, + "__regex__": TypeOrder.BSON_REGEX, +} + + # NOTE: This order is defined by the backend and cannot be changed. _TYPE_ORDER_MAP = { TypeOrder.NULL: 0, - TypeOrder.BOOLEAN: 1, - TypeOrder.NUMBER: 2, - TypeOrder.TIMESTAMP: 3, - TypeOrder.STRING: 4, - TypeOrder.BLOB: 5, - TypeOrder.REF: 6, - TypeOrder.GEO_POINT: 7, - TypeOrder.ARRAY: 8, - TypeOrder.VECTOR: 9, - TypeOrder.OBJECT: 10, + TypeOrder.BSON_MIN_KEY: 1, + TypeOrder.BOOLEAN: 2, + TypeOrder.NUMBER: 3, + TypeOrder.TIMESTAMP: 4, + TypeOrder.BSON_TIMESTAMP: 5, + TypeOrder.STRING: 6, + TypeOrder.BLOB: 7, + TypeOrder.BSON_BINARY: 8, + TypeOrder.REF: 9, + TypeOrder.BSON_OBJECT_ID: 10, + TypeOrder.GEO_POINT: 11, + TypeOrder.BSON_REGEX: 12, + TypeOrder.ARRAY: 13, + TypeOrder.VECTOR: 14, + TypeOrder.OBJECT: 15, + TypeOrder.BSON_MAX_KEY: 16, } @@ -102,22 +172,35 @@ def compare(cls, left, right) -> int: else: return 1 - if leftType == TypeOrder.NULL: - return 0 # nulls are all equal + if ( + leftType == TypeOrder.NULL + or leftType == TypeOrder.BSON_MIN_KEY + or leftType == TypeOrder.BSON_MAX_KEY + ): + return 0 # sentinels are equal elif leftType == TypeOrder.BOOLEAN: return cls._compare_to(left.boolean_value, right.boolean_value) elif leftType == TypeOrder.NUMBER: + # Handles int64, double, BSONInt32, and BSONDecimal128. return cls.compare_numbers(left, right) elif leftType == TypeOrder.TIMESTAMP: return cls.compare_timestamps(left, right) + elif leftType == TypeOrder.BSON_TIMESTAMP: + return cls.compare_bson_timestamps(left, right) elif leftType == TypeOrder.STRING: return cls._compare_to(left.string_value, right.string_value) elif leftType == TypeOrder.BLOB: return cls.compare_blobs(left, right) + elif leftType == TypeOrder.BSON_BINARY: + return cls.compare_bson_binaries(left, right) elif leftType == TypeOrder.REF: return cls.compare_resource_paths(left, right) + elif leftType == TypeOrder.BSON_OBJECT_ID: + return cls.compare_bson_object_ids(left, right) elif leftType == TypeOrder.GEO_POINT: return cls.compare_geo_points(left, right) + elif leftType == TypeOrder.BSON_REGEX: + return cls.compare_bson_regexes(left, right) elif leftType == TypeOrder.ARRAY: return cls.compare_arrays(left, right) elif leftType == TypeOrder.VECTOR: @@ -135,16 +218,76 @@ def compare_blobs(left, right) -> int: return Order._compare_to(left_bytes, right_bytes) + @staticmethod + def compare_bson_binaries(left, right) -> int: + l_bin = left.map_value.fields["__binary__"].bytes_value + r_bin = right.map_value.fields["__binary__"].bytes_value + + l_subtype = l_bin[0] if l_bin else 0 + r_subtype = r_bin[0] if r_bin else 0 + + cmp_subtype = Order._compare_to(l_subtype, r_subtype) + if cmp_subtype != 0: + return cmp_subtype + + return Order._compare_to( + l_bin[1:] if l_bin else b"", r_bin[1:] if r_bin else b"" + ) + + @staticmethod + def compare_bson_object_ids(left, right) -> int: + l_oid = left.map_value.fields["__oid__"].string_value + r_oid = right.map_value.fields["__oid__"].string_value + return Order._compare_to(l_oid, r_oid) + + @staticmethod + def compare_bson_regexes(left, right) -> int: + l_regex = left.map_value.fields["__regex__"].map_value.fields + r_regex = right.map_value.fields["__regex__"].map_value.fields + + l_pattern = l_regex["pattern"].string_value if "pattern" in l_regex else "" + r_pattern = r_regex["pattern"].string_value if "pattern" in r_regex else "" + cmp_pat = Order._compare_to(l_pattern, r_pattern) + if cmp_pat != 0: + return cmp_pat + + l_options = l_regex["options"].string_value if "options" in l_regex else "" + r_options = r_regex["options"].string_value if "options" in r_regex else "" + return Order._compare_to(l_options, r_options) + @staticmethod def compare_timestamps(left, right) -> Any: - left = left._pb.timestamp_value - right = right._pb.timestamp_value + left_pb = getattr(left, "_pb", left) + right_pb = getattr(right, "_pb", right) + + seconds = Order._compare_to( + left_pb.timestamp_value.seconds, right_pb.timestamp_value.seconds + ) + if seconds != 0: + return seconds - seconds = Order._compare_to(left.seconds or 0, right.seconds or 0) + return Order._compare_to( + left_pb.timestamp_value.nanos, right_pb.timestamp_value.nanos + ) + + @staticmethod + def compare_bson_timestamps(left, right) -> Any: + left_pb = getattr(left, "_pb", left) + right_pb = getattr(right, "_pb", right) + + l_ts = left_pb.map_value.fields["__request_timestamp__"].map_value.fields + l_sec = l_ts["seconds"].integer_value if "seconds" in l_ts else 0 + l_inc = l_ts["increment"].integer_value if "increment" in l_ts else 0 + + r_ts = right_pb.map_value.fields["__request_timestamp__"].map_value.fields + r_sec = r_ts["seconds"].integer_value if "seconds" in r_ts else 0 + r_inc = r_ts["increment"].integer_value if "increment" in r_ts else 0 + + seconds = Order._compare_to(l_sec, r_sec) if seconds != 0: return seconds - return Order._compare_to(left.nanos or 0, right.nanos or 0) + return Order._compare_to(l_inc, r_inc) @staticmethod def compare_geo_points(left, right) -> Any: @@ -231,9 +374,24 @@ def compare_objects(left, right) -> int: @staticmethod def compare_numbers(left, right) -> int: - left_value = decode_value(left, None) - right_value = decode_value(right, None) - return Order.compare_doubles(left_value, right_value) + """Compare numeric values across int, float, BSONInt32, and BSONDecimal128.""" + left_val = _to_number(left) + right_val = _to_number(right) + + left_nan = _is_nan(left_val) + right_nan = _is_nan(right_val) + if left_nan or right_nan: + return 0 if (left_nan and right_nan) else (-1 if left_nan else 1) + + # Python raises TypeError when comparing Decimal with float directly, + # but allows comparing Decimal with int. Convert float to Decimal + # to ensure safe cross-type comparison without float overflow. + if isinstance(left_val, decimal.Decimal) and isinstance(right_val, float): + right_val = decimal.Decimal(str(right_val)) + elif isinstance(right_val, decimal.Decimal) and isinstance(left_val, float): + left_val = decimal.Decimal(str(left_val)) + + return Order._compare_to(left_val, right_val) @staticmethod def compare_doubles(left, right) -> int: diff --git a/packages/google-cloud-firestore/google/cloud/firestore_v1/pipeline_result.py b/packages/google-cloud-firestore/google/cloud/firestore_v1/pipeline_result.py index e3fd74677a1e..eb25b1f6b24c 100644 --- a/packages/google-cloud-firestore/google/cloud/firestore_v1/pipeline_result.py +++ b/packages/google-cloud-firestore/google/cloud/firestore_v1/pipeline_result.py @@ -44,6 +44,7 @@ from google.cloud.firestore_v1.async_transaction import AsyncTransaction from google.cloud.firestore_v1.base_client import BaseClient from google.cloud.firestore_v1.base_document import BaseDocumentReference + from google.cloud.firestore_v1.bson import BSONType from google.cloud.firestore_v1.client import Client from google.cloud.firestore_v1.pipeline import Pipeline from google.cloud.firestore_v1.pipeline_expressions import Constant @@ -90,7 +91,7 @@ def __init__( self._update_time = update_time def __repr__(self): - return f"{type(self).__name__}(data={self.data()})" + return f"{type(self).__name__}(data={self.data()!r})" @property def ref(self) -> BaseDocumentReference | None: @@ -138,7 +139,7 @@ def __eq__(self, other: object) -> bool: return NotImplemented return (self._ref == other._ref) and (self._fields_pb == other._fields_pb) - def data(self) -> dict | "Vector" | None: + def data(self) -> dict | "Vector" | "BSONType" | bytes | None: """ Retrieves all fields in the result. diff --git a/packages/google-cloud-firestore/tests/system/test_system.py b/packages/google-cloud-firestore/tests/system/test_system.py index 3d95991aeb86..e5457102501a 100644 --- a/packages/google-cloud-firestore/tests/system/test_system.py +++ b/packages/google-cloud-firestore/tests/system/test_system.py @@ -1285,9 +1285,9 @@ def test_unicode_doc(client, cleanup, database): @pytest.mark.parametrize("database", [FIRESTORE_ENTERPRISE_DB], indirect=True) -def test_bson_document_writes(client, cleanup, database): - """Test write operations for BSON types on Enterprise DB.""" - collection_id = "bson_type_writes_" + UNIQUE_RESOURCE_ID +def test_bson_document_read_and_write(client, cleanup, database): + """Test read and write operations for BSON types on Enterprise DB.""" + collection_id = "bson_type_read_write_" + UNIQUE_RESOURCE_ID doc_ref = client.collection(collection_id).document("bson_doc") cleanup(doc_ref.delete) @@ -1296,6 +1296,7 @@ def test_bson_document_writes(client, cleanup, database): "min_key": BSONMinKey(), "max_key": BSONMaxKey(), "int32_val": BSONInt32(42), + "binary_val_sub0": b"hello", "binary_val_sub128": BSONBinary(b"world", subtype=128), "timestamp_val": BSONTimestamp(1700000000, 1), "regex_val": BSONRegex("^hello.*$", options="i"), @@ -1306,26 +1307,7 @@ def test_bson_document_writes(client, cleanup, database): snapshot = doc_ref.get() assert snapshot.exists - assert snapshot.to_dict() == { - "user_id": {"__oid__": "507f191e810c19729de860ea"}, - "min_key": {"__min__": None}, - "max_key": {"__max__": None}, - "int32_val": {"__int__": 42}, - "binary_val_sub128": {"__binary__": b"\x80world"}, - "timestamp_val": { - "__request_timestamp__": { - "seconds": 1700000000, - "increment": 1, - } - }, - "regex_val": { - "__regex__": { - "pattern": "^hello.*$", - "options": "i", - } - }, - "decimal128_val": {"__decimal128__": "123.45"}, - } + assert snapshot.to_dict() == bson_payload @pytest.mark.parametrize("database", [FIRESTORE_ENTERPRISE_DB], indirect=True) @@ -1362,12 +1344,34 @@ def test_bson_decimal128_special_values(client, cleanup, database): snapshot = doc_ref.get() assert snapshot.exists assert snapshot.to_dict() == { - "inf_val": {"__decimal128__": "Infinity"}, - "neg_inf_val": {"__decimal128__": "-Infinity"}, - "nan_val": {"__decimal128__": "NaN"}, + "inf_val": BSONDecimal128("Infinity"), + "neg_inf_val": BSONDecimal128("-Infinity"), + "nan_val": BSONDecimal128("NaN"), } +@pytest.mark.parametrize("database", [FIRESTORE_ENTERPRISE_DB], indirect=True) +def test_bson_query_ordering(client, cleanup, database): + """Test server query ordering for BSON types.""" + collection_id = "bson_ordering_" + UNIQUE_RESOURCE_ID + coll_ref = client.collection(collection_id) + + doc1 = coll_ref.document("doc1") + doc2 = coll_ref.document("doc2") + doc3 = coll_ref.document("doc3") + cleanup(doc1.delete) + cleanup(doc2.delete) + cleanup(doc3.delete) + + doc1.set({"val": BSONMinKey()}) + doc2.set({"val": BSONInt32(10)}) + doc3.set({"val": BSONMaxKey()}) + + query = coll_ref.order_by("val") + results = [doc.to_dict()["val"] for doc in query.stream()] + assert results == [BSONMinKey(), BSONInt32(10), BSONMaxKey()] + + @pytest.fixture(scope="module") def query_docs(client, database): collection_id = "qs" + UNIQUE_RESOURCE_ID diff --git a/packages/google-cloud-firestore/tests/system/test_system_async.py b/packages/google-cloud-firestore/tests/system/test_system_async.py index 824bed1b597b..27720fce3b23 100644 --- a/packages/google-cloud-firestore/tests/system/test_system_async.py +++ b/packages/google-cloud-firestore/tests/system/test_system_async.py @@ -1258,9 +1258,9 @@ async def test_list_collections_with_read_time(client, cleanup, database): @pytest.mark.asyncio @pytest.mark.parametrize("database", [FIRESTORE_ENTERPRISE_DB], indirect=True) -async def test_async_bson_document_writes(client, cleanup, database): - """Test async write operations for BSON types on Enterprise DB.""" - collection_id = "async_bson_type_writes_" + UNIQUE_RESOURCE_ID +async def test_async_bson_document_read_and_write(client, cleanup, database): + """Test async read and write operations for BSON types on Enterprise DB.""" + collection_id = "async_bson_type_read_write_" + UNIQUE_RESOURCE_ID doc_ref = client.collection(collection_id).document("bson_doc") cleanup(doc_ref.delete) @@ -1269,6 +1269,7 @@ async def test_async_bson_document_writes(client, cleanup, database): "min_key": BSONMinKey(), "max_key": BSONMaxKey(), "int32_val": BSONInt32(42), + "binary_val_sub0": b"hello", "binary_val_sub128": BSONBinary(b"world", subtype=128), "timestamp_val": BSONTimestamp(1700000000, 1), "regex_val": BSONRegex("^hello.*$", options="i"), @@ -1279,26 +1280,7 @@ async def test_async_bson_document_writes(client, cleanup, database): snapshot = await doc_ref.get() assert snapshot.exists - assert snapshot.to_dict() == { - "user_id": {"__oid__": "507f191e810c19729de860ea"}, - "min_key": {"__min__": None}, - "max_key": {"__max__": None}, - "int32_val": {"__int__": 42}, - "binary_val_sub128": {"__binary__": b"\x80world"}, - "timestamp_val": { - "__request_timestamp__": { - "seconds": 1700000000, - "increment": 1, - } - }, - "regex_val": { - "__regex__": { - "pattern": "^hello.*$", - "options": "i", - } - }, - "decimal128_val": {"__decimal128__": "123.45"}, - } + assert snapshot.to_dict() == bson_payload @pytest.mark.asyncio @@ -1337,9 +1319,9 @@ async def test_async_bson_decimal128_special_values(client, cleanup, database): snapshot = await doc_ref.get() assert snapshot.exists assert snapshot.to_dict() == { - "inf_val": {"__decimal128__": "Infinity"}, - "neg_inf_val": {"__decimal128__": "-Infinity"}, - "nan_val": {"__decimal128__": "NaN"}, + "inf_val": BSONDecimal128("Infinity"), + "neg_inf_val": BSONDecimal128("-Infinity"), + "nan_val": BSONDecimal128("NaN"), } diff --git a/packages/google-cloud-firestore/tests/unit/v1/test__helpers.py b/packages/google-cloud-firestore/tests/unit/v1/test__helpers.py index 4ce48424d3c4..88fe361eee31 100644 --- a/packages/google-cloud-firestore/tests/unit/v1/test__helpers.py +++ b/packages/google-cloud-firestore/tests/unit/v1/test__helpers.py @@ -706,6 +706,37 @@ def test_decode_dict_w_many_types(): assert decode_dict(value_fields, mock.sentinel.client) == expected +def test_decode_dict_w_bson_types(): + from google.cloud.firestore_v1._helpers import decode_dict, encode_dict + from google.cloud.firestore_v1.bson import ( + BSONBinary, + BSONDecimal128, + BSONInt32, + BSONMaxKey, + BSONMinKey, + BSONObjectId, + BSONRegex, + BSONTimestamp, + ) + + original_dict = { + "oid": BSONObjectId("507f191e810c19729de860ea"), + "min_k": BSONMinKey(), + "max_k": BSONMaxKey(), + "int32_v": BSONInt32(42), + "bin_sub0": b"hello", + "bin_sub0_empty": b"", + "bin_sub128": BSONBinary(b"world", subtype=128), + "ts_v": BSONTimestamp(1700000000, 1), + "regex_v": BSONRegex("^hello.*$", options="i"), + "dec_v": BSONDecimal128("123.45"), + } + + pb_fields = encode_dict(original_dict) + decoded = decode_dict(pb_fields, mock.sentinel.client) + assert decoded == original_dict + + def _dummy_ref_string(collection_id): from google.cloud.firestore_v1.base_client import DEFAULT_DATABASE diff --git a/packages/google-cloud-firestore/tests/unit/v1/test_bson.py b/packages/google-cloud-firestore/tests/unit/v1/test_bson.py index a07a9efa2e47..b5fdc3578278 100644 --- a/packages/google-cloud-firestore/tests/unit/v1/test_bson.py +++ b/packages/google-cloud-firestore/tests/unit/v1/test_bson.py @@ -31,18 +31,18 @@ BSONObjectId, BSONRegex, BSONTimestamp, - _BSONType, + BSONType, ) def test_bson_type_abc_cannot_be_instantiated(): with pytest.raises(TypeError): - _BSONType() # type: ignore + BSONType() # type: ignore def test_bson_type_inheritance(): oid = BSONObjectId("507f191e810c19729de860ea") - assert isinstance(oid, _BSONType) + assert isinstance(oid, BSONType) def test_bson_object_id_from_hex_string(): @@ -601,3 +601,13 @@ def test_bson_decimal128_copy(): def test_bson_decimal128_pickle(): d = BSONDecimal128("123.45") assert pickle.loads(pickle.dumps(d)) == d + + +def test_bson_from_dict_exception_fallback(): + # Corrupted or malformed BSON wire dictionary shapes gracefully return None + assert BSONType._from_dict({"__int__": "not-an-int"}) is None + assert BSONType._from_dict({"__oid__": "short"}) is None + assert BSONType._from_dict({"__decimal128__": "invalid-decimal"}) is None + assert BSONType._from_dict({"__unknown__": "value"}) is None + assert BSONType._from_dict("not-a-dict") is None + assert BSONType._from_dict({"a": 1, "b": 2}) is None diff --git a/packages/google-cloud-firestore/tests/unit/v1/test_order.py b/packages/google-cloud-firestore/tests/unit/v1/test_order.py index 1942a5298438..ff70eb441260 100644 --- a/packages/google-cloud-firestore/tests/unit/v1/test_order.py +++ b/packages/google-cloud-firestore/tests/unit/v1/test_order.py @@ -199,6 +199,100 @@ def test_order_all_value_present(): assert type_order in _TYPE_ORDER_MAP +def test_order_bson_type_ordering(): + from google.cloud.firestore_v1._helpers import encode_value + from google.cloud.firestore_v1.bson import ( + BSONBinary, + BSONDecimal128, + BSONInt32, + BSONMaxKey, + BSONMinKey, + BSONObjectId, + BSONRegex, + BSONTimestamp, + ) + from google.cloud.firestore_v1.order import Order + + min_k = encode_value(BSONMinKey()) + max_k = encode_value(BSONMaxKey()) + null_v = nullValue() + int32_v = encode_value(BSONInt32(10)) + int64_v = _int_value(10) + dec_v = encode_value(BSONDecimal128("10.0")) + ts_bson = encode_value(BSONTimestamp(100, 1)) + ts_native = _timestamp_value(100, 0) + bin_b = encode_value(BSONBinary(b"xyz", subtype=1)) + bytes_native = _blob_value(b"xyz") + ref_v = _reference_value("projects/p1/databases/d1/documents/c1/doc1") + oid_v = encode_value(BSONObjectId("507f191e810c19729de860ea")) + geo_v = _geoPoint_value(0, 0) + regex_v = encode_value(BSONRegex("abc")) + arr_v = _array_value() + map_v = _object_value({"a": 1}) + + # Test 16-rank ordering bounds + target = Order() + assert target.compare(null_v, min_k) == -1 + assert target.compare(min_k, null_v) == 1 + + assert target.compare(max_k, map_v) == 1 + assert target.compare(map_v, max_k) == -1 + + # Test numbers comparison equality across int32, int64, decimal128 + assert target.compare(int32_v, int64_v) == 0 + assert target.compare(int32_v, dec_v) == 0 + + # Test large decimal comparison exceeding float limit + large_dec = encode_value(BSONDecimal128("1e1000")) + assert target.compare(large_dec, _double_value(1e300)) == 1 + assert target.compare(_double_value(1e300), large_dec) == -1 + + # Test decimal NaN comparison + nan_dec = encode_value(BSONDecimal128("NaN")) + assert target.compare(nan_dec, int32_v) == -1 + assert target.compare(int32_v, nan_dec) == 1 + + # Test timestamp comparison (native timestamp < BSON timestamp with increment) + assert target.compare(ts_native, ts_bson) == -1 + assert target.compare(ts_bson, ts_native) == 1 + + # Test BSON timestamp comparison + ts_bson2 = encode_value(BSONTimestamp(100, 2)) + ts_bson_later = encode_value(BSONTimestamp(101, 0)) + assert target.compare(ts_bson, ts_bson2) == -1 + assert target.compare(ts_bson2, ts_bson) == 1 + assert target.compare(ts_bson, ts_bson_later) == -1 + assert target.compare(ts_bson_later, ts_bson) == 1 + assert target.compare(ts_bson, ts_bson) == 0 + + # Test BSON binary > bytes + assert target.compare(bytes_native, bin_b) == -1 + + # Test ObjectId rank (REF < OID < GEO_POINT) + assert target.compare(ref_v, oid_v) == -1 + assert target.compare(oid_v, geo_v) == -1 + + # Test Regex rank (GEO_POINT < REGEX < ARRAY) + assert target.compare(geo_v, regex_v) == -1 + assert target.compare(regex_v, arr_v) == -1 + + # Verify _BSON_KEY_TO_TYPE_ORDER mapping directly + from google.cloud.firestore_v1.order import _BSON_KEY_TO_TYPE_ORDER, TypeOrder + + expected_orders = { + "__min__": TypeOrder.BSON_MIN_KEY, + "__max__": TypeOrder.BSON_MAX_KEY, + "__oid__": TypeOrder.BSON_OBJECT_ID, + "__int__": TypeOrder.NUMBER, + "__decimal128__": TypeOrder.NUMBER, + "__binary__": TypeOrder.BSON_BINARY, + "__regex__": TypeOrder.BSON_REGEX, + "__request_timestamp__": TypeOrder.BSON_TIMESTAMP, + } + for key, expected_order in expected_orders.items(): + assert _BSON_KEY_TO_TYPE_ORDER.get(key) == expected_order + + def test_order_compare_w_objects_different_keys(): left = _object_value({"foo": 0}) right = _object_value({"bar": 0})