1212# See the License for the specific language governing permissions and
1313# limitations under the License.
1414
15+ import decimal
1516import math
1617from enum import Enum
1718from typing import Any
1819
1920from 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+
2262class 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 :
0 commit comments