1515# specific language governing permissions and limitations
1616# under the License.
1717# pylint:disable=redefined-outer-name
18- from concurrent .futures import ThreadPoolExecutor
19- from threading import Event
20- from typing import Any
21-
2218import pytest
2319
2420from pyiceberg .expressions import (
2521 AlwaysFalse ,
2622 AlwaysTrue ,
2723 And ,
24+ BooleanExpression ,
2825 EqualTo ,
2926 GreaterThan ,
3027 GreaterThanOrEqual ,
4542from pyiceberg .schema import Schema
4643from pyiceberg .transforms import DayTransform , IdentityTransform
4744from pyiceberg .typedef import Record
48- from pyiceberg .types import DoubleType , FloatType , IntegerType , NestedField , StringType , StructType , TimestampType
45+ from pyiceberg .types import DoubleType , FloatType , IntegerType , NestedField , StringType , TimestampType
4946
5047
5148def test_identity_transform_residual () -> None :
@@ -92,102 +89,25 @@ def test_identity_transform_residual() -> None:
9289 assert residual == AlwaysFalse ()
9390
9491
95- def test_residual_evaluator_does_not_mutate_prepared_state () -> None :
96- schema = Schema (NestedField (1 , "a" , IntegerType ()), NestedField (2 , "b" , IntegerType ()))
97- spec = PartitionSpec (
98- PartitionField (1 , 1001 , IdentityTransform (), "a_part" ),
99- PartitionField (2 , 1002 , IdentityTransform (), "b_part" ),
100- )
101- evaluator = residual_evaluator_of (
102- spec = spec ,
103- expr = And (EqualTo ("a" , 1 ), EqualTo ("b" , 1 )),
104- case_sensitive = True ,
105- schema = schema ,
106- )
107- initial_state = vars (evaluator ).copy ()
108-
109- assert evaluator .residual_for (Record (1 , 1 )) == AlwaysTrue ()
110- assert evaluator .residual_for (Record (0 , 0 )) == AlwaysFalse ()
111- assert evaluator .residual_for (Record (1 , 1 )) == AlwaysTrue ()
112-
113- assert isinstance (evaluator , ResidualVisitor )
114- assert vars (evaluator ) == initial_state
115-
116-
11792def test_residual_visitor_preserves_public_eval_api () -> None :
11893 schema = Schema (NestedField (1 , "a" , IntegerType ()))
11994 spec = PartitionSpec (PartitionField (1 , 1001 , IdentityTransform (), "a_part" ))
12095 visitor = ResidualVisitor (schema = schema , spec = spec , case_sensitive = True , expr = EqualTo ("a" , 1 ))
121- initial_state = vars (visitor ).copy ()
12296
12397 assert visitor .eval (Record (1 )) == AlwaysTrue ()
12498 assert visitor .eval (Record (0 )) == AlwaysFalse ()
125- assert vars (visitor ) == initial_state
126-
127-
128- def test_residual_evaluator_concurrent_calls_do_not_share_partitions () -> None :
129- class BlockingRecord (Record ):
130- def __init__ (self , first_read : Event , release_first_read : Event , * values : Any ) -> None :
131- super ().__init__ (* values )
132- self .first_read = first_read
133- self .release_first_read = release_first_read
134-
135- def __getitem__ (self , pos : int ) -> Any :
136- value = super ().__getitem__ (pos )
137- if pos == 0 :
138- self .first_read .set ()
139- if not self .release_first_read .wait (timeout = 5 ):
140- raise TimeoutError ("Timed out waiting to interleave residual evaluations" )
141- return value
142-
143- schema = Schema (NestedField (1 , "a" , IntegerType ()), NestedField (2 , "b" , IntegerType ()))
144- spec = PartitionSpec (
145- PartitionField (1 , 1001 , IdentityTransform (), "a_part" ),
146- PartitionField (2 , 1002 , IdentityTransform (), "b_part" ),
147- )
148- evaluator = residual_evaluator_of (
149- spec = spec ,
150- expr = And (EqualTo ("a" , 1 ), EqualTo ("b" , 1 )),
151- case_sensitive = True ,
152- schema = schema ,
153- )
154- first_read = Event ()
155- release_first_read = Event ()
15699
157- with ThreadPoolExecutor (max_workers = 2 ) as executor :
158- matching_result = executor .submit (
159- evaluator .residual_for ,
160- BlockingRecord (first_read , release_first_read , 1 , 1 ),
161- )
162- assert first_read .wait (timeout = 5 )
163100
164- try :
165- non_matching_result = executor . submit ( evaluator . residual_for , Record ( 0 , 0 )). result ( timeout = 5 )
166- finally :
167- release_first_read . set ()
101+ def test_residual_visitor_subclass_can_customize_evaluation () -> None :
102+ class FalseForTrueResidualVisitor ( ResidualVisitor ):
103+ def visit_true ( self ) -> BooleanExpression :
104+ return AlwaysFalse ()
168105
169- assert matching_result .result (timeout = 5 ) == AlwaysTrue ()
170- assert non_matching_result == AlwaysFalse ()
171-
172-
173- def test_partition_schema_reused_across_residuals (monkeypatch : pytest .MonkeyPatch ) -> None :
174- schema = Schema (NestedField (50 , "dateint" , IntegerType ()))
175- spec = PartitionSpec (PartitionField (50 , 1050 , IdentityTransform (), "dateint_part" ))
176- partition_type_calls = 0
177- original_partition_type = PartitionSpec .partition_type
178-
179- def counting_partition_type (self : PartitionSpec , schema : Schema ) -> StructType :
180- nonlocal partition_type_calls
181- partition_type_calls += 1
182- return original_partition_type (self , schema )
183-
184- monkeypatch .setattr (PartitionSpec , "partition_type" , counting_partition_type )
185-
186- evaluator = residual_evaluator_of (spec = spec , expr = EqualTo ("dateint" , 20170815 ), case_sensitive = True , schema = schema )
106+ schema = Schema (NestedField (1 , "a" , IntegerType ()))
107+ spec = PartitionSpec (PartitionField (1 , 1001 , IdentityTransform (), "a_part" ))
108+ visitor = FalseForTrueResidualVisitor (schema = schema , spec = spec , case_sensitive = True , expr = AlwaysTrue ())
187109
188- assert evaluator .residual_for (Record (20170815 )) == AlwaysTrue ()
189- assert evaluator .residual_for (Record (20170816 )) == AlwaysFalse ()
190- assert partition_type_calls == 1
110+ assert visitor .eval (Record (1 )) == AlwaysFalse ()
191111
192112
193113def test_case_insensitive_identity_transform_residuals () -> None :
@@ -313,6 +233,9 @@ def test_is_not_nan() -> None:
313233 res_eval = residual_evaluator_of (spec = spec , expr = predicate , case_sensitive = True , schema = schema )
314234
315235 residual = res_eval .residual_for (Record (None ))
236+ assert residual == AlwaysTrue ()
237+
238+ residual = res_eval .residual_for (Record (float ("nan" )))
316239 assert residual == AlwaysFalse ()
317240
318241 residual = res_eval .residual_for (Record (2 ))
@@ -325,6 +248,9 @@ def test_is_not_nan() -> None:
325248 res_eval = residual_evaluator_of (spec = spec , expr = predicate , case_sensitive = True , schema = schema )
326249
327250 residual = res_eval .residual_for (Record (None ))
251+ assert residual == AlwaysTrue ()
252+
253+ residual = res_eval .residual_for (Record (float ("nan" )))
328254 assert residual == AlwaysFalse ()
329255
330256 residual = res_eval .residual_for (Record (2 ))
0 commit comments