Skip to content

Commit 3250d05

Browse files
authored
feat(firestore): add BSON cross-type query ordering support
1 parent 6291c61 commit 3250d05

3 files changed

Lines changed: 193 additions & 24 deletions

File tree

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

Lines changed: 116 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -36,10 +36,16 @@ class TypeOrder(Enum):
3636
ARRAY = 8
3737
OBJECT = 9
3838
VECTOR = 10
39+
BSON_MIN_KEY = 11
40+
BSON_MAX_KEY = 12
41+
BSON_OBJECT_ID = 13
42+
BSON_BINARY = 14
43+
BSON_REGEX = 15
3944

4045
@staticmethod
4146
def from_value(value) -> Any:
42-
v = value._pb.WhichOneof("value_type")
47+
value_pb = getattr(value, "_pb", value)
48+
v = value_pb.WhichOneof("value_type")
4349
lut = {
4450
"null_value": TypeOrder.NULL,
4551
"boolean_value": TypeOrder.BOOLEAN,
@@ -58,27 +64,46 @@ def from_value(value) -> Any:
5864
raise ValueError(f"Could not detect value type for {v}")
5965

6066
if v == "map_value":
61-
if (
62-
"__type__" in value.map_value.fields
63-
and value.map_value.fields["__type__"].string_value == "__vector__"
64-
):
67+
fields = value_pb.map_value.fields
68+
if len(fields) == 1:
69+
key = next(iter(fields))
70+
if key == "__min__":
71+
return TypeOrder.BSON_MIN_KEY
72+
if key == "__max__":
73+
return TypeOrder.BSON_MAX_KEY
74+
if key == "__oid__":
75+
return TypeOrder.BSON_OBJECT_ID
76+
if key in ("__int__", "__decimal128__"):
77+
return TypeOrder.NUMBER
78+
if key == "__binary__":
79+
return TypeOrder.BSON_BINARY
80+
if key == "__regex__":
81+
return TypeOrder.BSON_REGEX
82+
if key == "__request_timestamp__":
83+
return TypeOrder.TIMESTAMP
84+
if "__type__" in fields and fields["__type__"].string_value == "__vector__":
6585
return TypeOrder.VECTOR
6686
return lut[v]
6787

6888

6989
# NOTE: This order is defined by the backend and cannot be changed.
7090
_TYPE_ORDER_MAP = {
7191
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,
92+
TypeOrder.BSON_MIN_KEY: 1,
93+
TypeOrder.BOOLEAN: 2,
94+
TypeOrder.NUMBER: 3,
95+
TypeOrder.TIMESTAMP: 4,
96+
TypeOrder.STRING: 5,
97+
TypeOrder.BLOB: 6,
98+
TypeOrder.BSON_BINARY: 7,
99+
TypeOrder.REF: 8,
100+
TypeOrder.BSON_OBJECT_ID: 9,
101+
TypeOrder.GEO_POINT: 10,
102+
TypeOrder.BSON_REGEX: 11,
103+
TypeOrder.ARRAY: 12,
104+
TypeOrder.VECTOR: 13,
105+
TypeOrder.OBJECT: 14,
106+
TypeOrder.BSON_MAX_KEY: 15,
82107
}
83108

84109

@@ -102,8 +127,12 @@ def compare(cls, left, right) -> int:
102127
else:
103128
return 1
104129

105-
if leftType == TypeOrder.NULL:
106-
return 0 # nulls are all equal
130+
if (
131+
leftType == TypeOrder.NULL
132+
or leftType == TypeOrder.BSON_MIN_KEY
133+
or leftType == TypeOrder.BSON_MAX_KEY
134+
):
135+
return 0 # sentinels are equal
107136
elif leftType == TypeOrder.BOOLEAN:
108137
return cls._compare_to(left.boolean_value, right.boolean_value)
109138
elif leftType == TypeOrder.NUMBER:
@@ -114,10 +143,16 @@ def compare(cls, left, right) -> int:
114143
return cls._compare_to(left.string_value, right.string_value)
115144
elif leftType == TypeOrder.BLOB:
116145
return cls.compare_blobs(left, right)
146+
elif leftType == TypeOrder.BSON_BINARY:
147+
return cls.compare_bson_binaries(left, right)
117148
elif leftType == TypeOrder.REF:
118149
return cls.compare_resource_paths(left, right)
150+
elif leftType == TypeOrder.BSON_OBJECT_ID:
151+
return cls.compare_bson_object_ids(left, right)
119152
elif leftType == TypeOrder.GEO_POINT:
120153
return cls.compare_geo_points(left, right)
154+
elif leftType == TypeOrder.BSON_REGEX:
155+
return cls.compare_bson_regexes(left, right)
121156
elif leftType == TypeOrder.ARRAY:
122157
return cls.compare_arrays(left, right)
123158
elif leftType == TypeOrder.VECTOR:
@@ -135,16 +170,69 @@ def compare_blobs(left, right) -> int:
135170

136171
return Order._compare_to(left_bytes, right_bytes)
137172

173+
@staticmethod
174+
def compare_bson_binaries(left, right) -> int:
175+
l_bin = left.map_value.fields["__binary__"].bytes_value
176+
r_bin = right.map_value.fields["__binary__"].bytes_value
177+
178+
l_subtype = l_bin[0] if l_bin else 0
179+
r_subtype = r_bin[0] if r_bin else 0
180+
181+
cmp_subtype = Order._compare_to(l_subtype, r_subtype)
182+
if cmp_subtype != 0:
183+
return cmp_subtype
184+
185+
return Order._compare_to(
186+
l_bin[1:] if l_bin else b"", r_bin[1:] if r_bin else b""
187+
)
188+
189+
@staticmethod
190+
def compare_bson_object_ids(left, right) -> int:
191+
l_oid = left.map_value.fields["__oid__"].string_value
192+
r_oid = right.map_value.fields["__oid__"].string_value
193+
return Order._compare_to(l_oid, r_oid)
194+
195+
@staticmethod
196+
def compare_bson_regexes(left, right) -> int:
197+
l_regex = left.map_value.fields["__regex__"].map_value.fields
198+
r_regex = right.map_value.fields["__regex__"].map_value.fields
199+
200+
l_pattern = l_regex["pattern"].string_value if "pattern" in l_regex else ""
201+
r_pattern = r_regex["pattern"].string_value if "pattern" in r_regex else ""
202+
cmp_pat = Order._compare_to(l_pattern, r_pattern)
203+
if cmp_pat != 0:
204+
return cmp_pat
205+
206+
l_options = l_regex["options"].string_value if "options" in l_regex else ""
207+
r_options = r_regex["options"].string_value if "options" in r_regex else ""
208+
return Order._compare_to(l_options, r_options)
209+
138210
@staticmethod
139211
def compare_timestamps(left, right) -> Any:
140-
left = left._pb.timestamp_value
141-
right = right._pb.timestamp_value
212+
left_pb = getattr(left, "_pb", left)
213+
right_pb = getattr(right, "_pb", right)
214+
215+
if left_pb.WhichOneof("value_type") == "map_value":
216+
l_ts = left_pb.map_value.fields["__request_timestamp__"].map_value.fields
217+
l_sec = l_ts["seconds"].integer_value if "seconds" in l_ts else 0
218+
l_inc = l_ts["increment"].integer_value if "increment" in l_ts else 0
219+
else:
220+
l_sec = left_pb.timestamp_value.seconds
221+
l_inc = left_pb.timestamp_value.nanos
142222

143-
seconds = Order._compare_to(left.seconds or 0, right.seconds or 0)
223+
if right_pb.WhichOneof("value_type") == "map_value":
224+
r_ts = right_pb.map_value.fields["__request_timestamp__"].map_value.fields
225+
r_sec = r_ts["seconds"].integer_value if "seconds" in r_ts else 0
226+
r_inc = r_ts["increment"].integer_value if "increment" in r_ts else 0
227+
else:
228+
r_sec = right_pb.timestamp_value.seconds
229+
r_inc = right_pb.timestamp_value.nanos
230+
231+
seconds = Order._compare_to(l_sec, r_sec)
144232
if seconds != 0:
145233
return seconds
146234

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

149237
@staticmethod
150238
def compare_geo_points(left, right) -> Any:
@@ -231,9 +319,13 @@ def compare_objects(left, right) -> int:
231319

232320
@staticmethod
233321
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)
322+
left_val = decode_value(left, None, decode_bson=True)
323+
right_val = decode_value(right, None, decode_bson=True)
324+
if hasattr(left_val, "value"):
325+
left_val = left_val.value
326+
if hasattr(right_val, "value"):
327+
right_val = right_val.value
328+
return Order.compare_doubles(float(left_val), float(right_val))
237329

238330
@staticmethod
239331
def compare_doubles(left, right) -> int:

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

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1310,6 +1310,25 @@ def test_bson_document_read_and_write(client, cleanup, database):
13101310
assert snapshot.to_dict(decode_bson=True) == bson_payload
13111311

13121312

1313+
def test_bson_query_ordering(client, cleanup, database):
1314+
"""Test server query ordering for BSON types."""
1315+
collection_id = "bson_ordering_" + UNIQUE_RESOURCE_ID
1316+
coll_ref = client.collection(collection_id)
1317+
1318+
doc1 = coll_ref.document("doc1")
1319+
doc2 = coll_ref.document("doc2")
1320+
doc3 = coll_ref.document("doc3")
1321+
cleanup.extend([doc1.delete, doc2.delete, doc3.delete])
1322+
1323+
doc1.set({"val": BSONMinKey()})
1324+
doc2.set({"val": BSONInt32(10)})
1325+
doc3.set({"val": BSONMaxKey()})
1326+
1327+
query = coll_ref.order_by("val")
1328+
results = [doc.to_dict(decode_bson=True)["val"] for doc in query.stream()]
1329+
assert results == [BSONMinKey(), BSONInt32(10), BSONMaxKey()]
1330+
1331+
13131332
@pytest.fixture(scope="module")
13141333
def query_docs(client, database):
13151334
collection_id = "qs" + UNIQUE_RESOURCE_ID

packages/google-cloud-firestore/tests/unit/v1/test_order.py

Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -199,6 +199,64 @@ def test_order_all_value_present():
199199
assert type_order in _TYPE_ORDER_MAP
200200

201201

202+
def test_order_bson_type_ordering():
203+
from google.cloud.firestore_v1._helpers import encode_value
204+
from google.cloud.firestore_v1.bson import (
205+
BSONBinary,
206+
BSONDecimal128,
207+
BSONInt32,
208+
BSONMaxKey,
209+
BSONMinKey,
210+
BSONObjectId,
211+
BSONRegex,
212+
BSONTimestamp,
213+
)
214+
from google.cloud.firestore_v1.order import Order
215+
216+
min_k = encode_value(BSONMinKey())
217+
max_k = encode_value(BSONMaxKey())
218+
null_v = nullValue()
219+
int32_v = encode_value(BSONInt32(10))
220+
int64_v = _int_value(10)
221+
dec_v = encode_value(BSONDecimal128("10.0"))
222+
ts_bson = encode_value(BSONTimestamp(100, 1))
223+
ts_native = _timestamp_value(100, 0)
224+
bin_b = encode_value(BSONBinary(b"xyz", subtype=1))
225+
bytes_native = _blob_value(b"xyz")
226+
ref_v = _reference_value("projects/p1/databases/d1/documents/c1/doc1")
227+
oid_v = encode_value(BSONObjectId("507f191e810c19729de860ea"))
228+
geo_v = _geoPoint_value(0, 0)
229+
regex_v = encode_value(BSONRegex("abc"))
230+
arr_v = _array_value()
231+
map_v = _object_value({"a": 1})
232+
233+
# Test 16-rank ordering bounds
234+
target = Order()
235+
assert target.compare(null_v, min_k) == -1
236+
assert target.compare(min_k, null_v) == 1
237+
238+
assert target.compare(max_k, map_v) == 1
239+
assert target.compare(map_v, max_k) == -1
240+
241+
# Test numbers comparison equality across int32, int64, decimal128
242+
assert target.compare(int32_v, int64_v) == 0
243+
assert target.compare(int32_v, dec_v) == 0
244+
245+
# Test timestamp comparison (native timestamp < BSON timestamp with increment)
246+
assert target.compare(ts_native, ts_bson) == -1
247+
248+
# Test BSON binary > bytes
249+
assert target.compare(bytes_native, bin_b) == -1
250+
251+
# Test ObjectId rank (REF < OID < GEO_POINT)
252+
assert target.compare(ref_v, oid_v) == -1
253+
assert target.compare(oid_v, geo_v) == -1
254+
255+
# Test Regex rank (GEO_POINT < REGEX < ARRAY)
256+
assert target.compare(geo_v, regex_v) == -1
257+
assert target.compare(regex_v, arr_v) == -1
258+
259+
202260
def test_order_compare_w_objects_different_keys():
203261
left = _object_value({"foo": 0})
204262
right = _object_value({"bar": 0})

0 commit comments

Comments
 (0)