Skip to content

Commit 13be295

Browse files
authored
feat(firestore): add BSON cross-type query ordering support (#18405)
## Context & Problem Firestore queries require client-side and cross-type value ordering support for BSON types, including BSONMinKey, BSONMaxKey, BSONObjectId, BSONInt32, BSONDecimal128, BSONBinary, BSONRegex, and BSONTimestamp, aligning cross-type ordering with backend specifications. ## Summary of Changes - Add `_BSON_KEY_TO_TYPE_ORDER` mapping dictionary in `order.py` for $O(1)$ wire-key resolution to `TypeOrder`. - Consolidate cross-type numeric comparisons in `Order.compare_numbers` with inline unboxing, unified NaN handling, and safe float-to-`Decimal` conversion to prevent float overflow and precision loss. - Restore `Order.compare_doubles` to standard float comparisons. - Implement cross-type timestamp comparison in `Order.compare_timestamps` supporting native Firestore timestamps and `BSONTimestamp`. - Add unit tests in `test_order.py` verifying BSON type ordering, `1e1000` large Decimal comparisons, Decimal `NaN` handling, and wire-key mapping lookups. - Update system test to verify query ordering across BSON types. ## Verification - Ran local unit test suite: `pytest tests/unit/v1/test_order.py tests/unit/v1/test_bson.py` (112 passed in 1.18s). - Ran full test suite across Python 3.11 with python and upb protobuf implementations (3,708 passed). - Ran linter suite: `nox -e lint` (all checks passed, 276 files formatted). - Ran type checker: `nox -s mypy-3.11` (Success: no issues found in 111 source files).
1 parent ce2544f commit 13be295

3 files changed

Lines changed: 298 additions & 24 deletions

File tree

‎packages/google-cloud-firestore/google/cloud/firestore_v1/order.py‎

Lines changed: 182 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -12,13 +12,53 @@
1212
# See the License for the specific language governing permissions and
1313
# limitations under the License.
1414

15+
import decimal
1516
import math
1617
from enum import Enum
1718
from typing import Any
1819

1920
from google.cloud.firestore_v1._helpers import GeoPoint, decode_value
2021

2122

23+
def _to_number(val: Any) -> Any:
24+
"""Extract a numeric value (int, float, Decimal) from a Value protobuf or Python value.
25+
26+
Directly inspects the protobuf value_type without calling decode_value()
27+
for optimal performance.
28+
"""
29+
value_pb = getattr(val, "_pb", val)
30+
which = (
31+
value_pb.WhichOneof("value_type") if hasattr(value_pb, "WhichOneof") else None
32+
)
33+
34+
if which == "integer_value":
35+
return value_pb.integer_value
36+
elif which == "double_value":
37+
return value_pb.double_value
38+
elif which == "map_value":
39+
fields = value_pb.map_value.fields
40+
if "__int__" in fields:
41+
return fields["__int__"].integer_value
42+
elif "__decimal128__" in fields:
43+
return decimal.Decimal(fields["__decimal128__"].string_value)
44+
45+
num = decode_value(val, None)
46+
to_decimal = getattr(num, "to_decimal", None)
47+
return to_decimal() if callable(to_decimal) else getattr(num, "value", num)
48+
49+
50+
def _is_nan(val: Any) -> bool:
51+
"""Check if a numeric value is NaN, safely handling OverflowError and non-floats."""
52+
if hasattr(val, "is_nan"):
53+
return val.is_nan()
54+
if isinstance(val, (int, decimal.Decimal)):
55+
return False
56+
try:
57+
return math.isnan(val)
58+
except (TypeError, OverflowError):
59+
return False
60+
61+
2262
class TypeOrder(Enum):
2363
"""The supported Data Type.
2464
@@ -36,10 +76,17 @@ class TypeOrder(Enum):
3676
ARRAY = 8
3777
OBJECT = 9
3878
VECTOR = 10
79+
BSON_MIN_KEY = 11
80+
BSON_MAX_KEY = 12
81+
BSON_OBJECT_ID = 13
82+
BSON_BINARY = 14
83+
BSON_REGEX = 15
84+
BSON_TIMESTAMP = 16
3985

4086
@staticmethod
4187
def from_value(value) -> Any:
42-
v = value._pb.WhichOneof("value_type")
88+
value_pb = getattr(value, "_pb", value)
89+
v = value_pb.WhichOneof("value_type")
4390
lut = {
4491
"null_value": TypeOrder.NULL,
4592
"boolean_value": TypeOrder.BOOLEAN,
@@ -58,27 +105,50 @@ def from_value(value) -> Any:
58105
raise ValueError(f"Could not detect value type for {v}")
59106

60107
if v == "map_value":
61-
if (
62-
"__type__" in value.map_value.fields
63-
and value.map_value.fields["__type__"].string_value == "__vector__"
64-
):
108+
fields = value_pb.map_value.fields
109+
if len(fields) == 1:
110+
key = next(iter(fields))
111+
bson_order = _BSON_KEY_TO_TYPE_ORDER.get(key)
112+
if bson_order is not None:
113+
return bson_order
114+
if "__type__" in fields and fields["__type__"].string_value == "__vector__":
65115
return TypeOrder.VECTOR
66116
return lut[v]
67117

68118

119+
# Maps BSON wire map keys directly to their corresponding TypeOrder.
120+
# BSONInt32 and BSONDecimal128 map to TypeOrder.NUMBER, enabling cross-type comparisons.
121+
_BSON_KEY_TO_TYPE_ORDER = {
122+
"__min__": TypeOrder.BSON_MIN_KEY,
123+
"__max__": TypeOrder.BSON_MAX_KEY,
124+
"__oid__": TypeOrder.BSON_OBJECT_ID,
125+
"__int__": TypeOrder.NUMBER,
126+
"__decimal128__": TypeOrder.NUMBER,
127+
"__binary__": TypeOrder.BSON_BINARY,
128+
"__request_timestamp__": TypeOrder.BSON_TIMESTAMP,
129+
"__regex__": TypeOrder.BSON_REGEX,
130+
}
131+
132+
69133
# NOTE: This order is defined by the backend and cannot be changed.
70134
_TYPE_ORDER_MAP = {
71135
TypeOrder.NULL: 0,
72-
TypeOrder.BOOLEAN: 1,
73-
TypeOrder.NUMBER: 2,
74-
TypeOrder.TIMESTAMP: 3,
75-
TypeOrder.STRING: 4,
76-
TypeOrder.BLOB: 5,
77-
TypeOrder.REF: 6,
78-
TypeOrder.GEO_POINT: 7,
79-
TypeOrder.ARRAY: 8,
80-
TypeOrder.VECTOR: 9,
81-
TypeOrder.OBJECT: 10,
136+
TypeOrder.BSON_MIN_KEY: 1,
137+
TypeOrder.BOOLEAN: 2,
138+
TypeOrder.NUMBER: 3,
139+
TypeOrder.TIMESTAMP: 4,
140+
TypeOrder.BSON_TIMESTAMP: 5,
141+
TypeOrder.STRING: 6,
142+
TypeOrder.BLOB: 7,
143+
TypeOrder.BSON_BINARY: 8,
144+
TypeOrder.REF: 9,
145+
TypeOrder.BSON_OBJECT_ID: 10,
146+
TypeOrder.GEO_POINT: 11,
147+
TypeOrder.BSON_REGEX: 12,
148+
TypeOrder.ARRAY: 13,
149+
TypeOrder.VECTOR: 14,
150+
TypeOrder.OBJECT: 15,
151+
TypeOrder.BSON_MAX_KEY: 16,
82152
}
83153

84154

@@ -102,22 +172,35 @@ def compare(cls, left, right) -> int:
102172
else:
103173
return 1
104174

105-
if leftType == TypeOrder.NULL:
106-
return 0 # nulls are all equal
175+
if (
176+
leftType == TypeOrder.NULL
177+
or leftType == TypeOrder.BSON_MIN_KEY
178+
or leftType == TypeOrder.BSON_MAX_KEY
179+
):
180+
return 0 # sentinels are equal
107181
elif leftType == TypeOrder.BOOLEAN:
108182
return cls._compare_to(left.boolean_value, right.boolean_value)
109183
elif leftType == TypeOrder.NUMBER:
184+
# Handles int64, double, BSONInt32, and BSONDecimal128.
110185
return cls.compare_numbers(left, right)
111186
elif leftType == TypeOrder.TIMESTAMP:
112187
return cls.compare_timestamps(left, right)
188+
elif leftType == TypeOrder.BSON_TIMESTAMP:
189+
return cls.compare_bson_timestamps(left, right)
113190
elif leftType == TypeOrder.STRING:
114191
return cls._compare_to(left.string_value, right.string_value)
115192
elif leftType == TypeOrder.BLOB:
116193
return cls.compare_blobs(left, right)
194+
elif leftType == TypeOrder.BSON_BINARY:
195+
return cls.compare_bson_binaries(left, right)
117196
elif leftType == TypeOrder.REF:
118197
return cls.compare_resource_paths(left, right)
198+
elif leftType == TypeOrder.BSON_OBJECT_ID:
199+
return cls.compare_bson_object_ids(left, right)
119200
elif leftType == TypeOrder.GEO_POINT:
120201
return cls.compare_geo_points(left, right)
202+
elif leftType == TypeOrder.BSON_REGEX:
203+
return cls.compare_bson_regexes(left, right)
121204
elif leftType == TypeOrder.ARRAY:
122205
return cls.compare_arrays(left, right)
123206
elif leftType == TypeOrder.VECTOR:
@@ -135,16 +218,76 @@ def compare_blobs(left, right) -> int:
135218

136219
return Order._compare_to(left_bytes, right_bytes)
137220

221+
@staticmethod
222+
def compare_bson_binaries(left, right) -> int:
223+
l_bin = left.map_value.fields["__binary__"].bytes_value
224+
r_bin = right.map_value.fields["__binary__"].bytes_value
225+
226+
l_subtype = l_bin[0] if l_bin else 0
227+
r_subtype = r_bin[0] if r_bin else 0
228+
229+
cmp_subtype = Order._compare_to(l_subtype, r_subtype)
230+
if cmp_subtype != 0:
231+
return cmp_subtype
232+
233+
return Order._compare_to(
234+
l_bin[1:] if l_bin else b"", r_bin[1:] if r_bin else b""
235+
)
236+
237+
@staticmethod
238+
def compare_bson_object_ids(left, right) -> int:
239+
l_oid = left.map_value.fields["__oid__"].string_value
240+
r_oid = right.map_value.fields["__oid__"].string_value
241+
return Order._compare_to(l_oid, r_oid)
242+
243+
@staticmethod
244+
def compare_bson_regexes(left, right) -> int:
245+
l_regex = left.map_value.fields["__regex__"].map_value.fields
246+
r_regex = right.map_value.fields["__regex__"].map_value.fields
247+
248+
l_pattern = l_regex["pattern"].string_value if "pattern" in l_regex else ""
249+
r_pattern = r_regex["pattern"].string_value if "pattern" in r_regex else ""
250+
cmp_pat = Order._compare_to(l_pattern, r_pattern)
251+
if cmp_pat != 0:
252+
return cmp_pat
253+
254+
l_options = l_regex["options"].string_value if "options" in l_regex else ""
255+
r_options = r_regex["options"].string_value if "options" in r_regex else ""
256+
return Order._compare_to(l_options, r_options)
257+
138258
@staticmethod
139259
def compare_timestamps(left, right) -> Any:
140-
left = left._pb.timestamp_value
141-
right = right._pb.timestamp_value
260+
left_pb = getattr(left, "_pb", left)
261+
right_pb = getattr(right, "_pb", right)
262+
263+
seconds = Order._compare_to(
264+
left_pb.timestamp_value.seconds, right_pb.timestamp_value.seconds
265+
)
266+
if seconds != 0:
267+
return seconds
142268

143-
seconds = Order._compare_to(left.seconds or 0, right.seconds or 0)
269+
return Order._compare_to(
270+
left_pb.timestamp_value.nanos, right_pb.timestamp_value.nanos
271+
)
272+
273+
@staticmethod
274+
def compare_bson_timestamps(left, right) -> Any:
275+
left_pb = getattr(left, "_pb", left)
276+
right_pb = getattr(right, "_pb", right)
277+
278+
l_ts = left_pb.map_value.fields["__request_timestamp__"].map_value.fields
279+
l_sec = l_ts["seconds"].integer_value if "seconds" in l_ts else 0
280+
l_inc = l_ts["increment"].integer_value if "increment" in l_ts else 0
281+
282+
r_ts = right_pb.map_value.fields["__request_timestamp__"].map_value.fields
283+
r_sec = r_ts["seconds"].integer_value if "seconds" in r_ts else 0
284+
r_inc = r_ts["increment"].integer_value if "increment" in r_ts else 0
285+
286+
seconds = Order._compare_to(l_sec, r_sec)
144287
if seconds != 0:
145288
return seconds
146289

147-
return Order._compare_to(left.nanos or 0, right.nanos or 0)
290+
return Order._compare_to(l_inc, r_inc)
148291

149292
@staticmethod
150293
def compare_geo_points(left, right) -> Any:
@@ -231,9 +374,24 @@ def compare_objects(left, right) -> int:
231374

232375
@staticmethod
233376
def compare_numbers(left, right) -> int:
234-
left_value = decode_value(left, None)
235-
right_value = decode_value(right, None)
236-
return Order.compare_doubles(left_value, right_value)
377+
"""Compare numeric values across int, float, BSONInt32, and BSONDecimal128."""
378+
left_val = _to_number(left)
379+
right_val = _to_number(right)
380+
381+
left_nan = _is_nan(left_val)
382+
right_nan = _is_nan(right_val)
383+
if left_nan or right_nan:
384+
return 0 if (left_nan and right_nan) else (-1 if left_nan else 1)
385+
386+
# Python raises TypeError when comparing Decimal with float directly,
387+
# but allows comparing Decimal with int. Convert float to Decimal
388+
# to ensure safe cross-type comparison without float overflow.
389+
if isinstance(left_val, decimal.Decimal) and isinstance(right_val, float):
390+
right_val = decimal.Decimal(str(right_val))
391+
elif isinstance(right_val, decimal.Decimal) and isinstance(left_val, float):
392+
left_val = decimal.Decimal(str(left_val))
393+
394+
return Order._compare_to(left_val, right_val)
237395

238396
@staticmethod
239397
def compare_doubles(left, right) -> int:

‎packages/google-cloud-firestore/tests/system/test_system.py‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1350,6 +1350,28 @@ def test_bson_decimal128_special_values(client, cleanup, database):
13501350
}
13511351

13521352

1353+
@pytest.mark.parametrize("database", [FIRESTORE_ENTERPRISE_DB], indirect=True)
1354+
def test_bson_query_ordering(client, cleanup, database):
1355+
"""Test server query ordering for BSON types."""
1356+
collection_id = "bson_ordering_" + UNIQUE_RESOURCE_ID
1357+
coll_ref = client.collection(collection_id)
1358+
1359+
doc1 = coll_ref.document("doc1")
1360+
doc2 = coll_ref.document("doc2")
1361+
doc3 = coll_ref.document("doc3")
1362+
cleanup(doc1.delete)
1363+
cleanup(doc2.delete)
1364+
cleanup(doc3.delete)
1365+
1366+
doc1.set({"val": BSONMinKey()})
1367+
doc2.set({"val": BSONInt32(10)})
1368+
doc3.set({"val": BSONMaxKey()})
1369+
1370+
query = coll_ref.order_by("val")
1371+
results = [doc.to_dict()["val"] for doc in query.stream()]
1372+
assert results == [BSONMinKey(), BSONInt32(10), BSONMaxKey()]
1373+
1374+
13531375
@pytest.fixture(scope="module")
13541376
def query_docs(client, database):
13551377
collection_id = "qs" + UNIQUE_RESOURCE_ID

0 commit comments

Comments
 (0)