@@ -579,6 +579,40 @@ def f(**kwargs):
579579 with self .assertRaises (RuntimeError ):
580580 result = p (** {BadStr ("poison" ): "new_value" })
581581
582+ def test_call_safety_against_reentrant_mutation (self ):
583+ def old_function (* args , ** kwargs ):
584+ return "old_function" , args , kwargs
585+
586+ def new_function (* args , ** kwargs ):
587+ return "new_function" , args , kwargs
588+
589+ g_partial = None
590+
591+ class EvilKey (str ):
592+ armed = False
593+ def __hash__ (self ):
594+ if EvilKey .armed and g_partial is not None :
595+ EvilKey .armed = False
596+ new_args_tuple = ("new_arg" ,)
597+ new_keywords_dict = {"new_keyword" : None }
598+ new_tuple_state = (new_function , new_args_tuple , new_keywords_dict , None )
599+ g_partial .__setstate__ (new_tuple_state )
600+ gc .collect ()
601+ return str .__hash__ (self )
602+
603+ g_partial = functools .partial (old_function , "old_arg" , old_keyword = None )
604+
605+ kwargs = {EvilKey ("evil_key" ): None }
606+ EvilKey .armed = True
607+
608+ result = g_partial (** kwargs )
609+ expected = ("old_function" , ("old_arg" ,), {"old_keyword" : None , "evil_key" : None })
610+ self .assertEqual (result , expected )
611+
612+ result = g_partial ()
613+ expected = ("new_function" , ("new_arg" ,), {"new_keyword" : None })
614+ self .assertEqual (result , expected )
615+
582616@unittest .skipUnless (c_functools , 'requires the C _functools module' )
583617class TestPartialC (TestPartial , unittest .TestCase ):
584618 if c_functools :
0 commit comments