@@ -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 :
0 commit comments