@@ -267,6 +267,75 @@ def test_access_via_instance(self):
267267 self .assertEqual (info .misses , 1 )
268268
269269
270+ class OverrideBase :
271+ def __init__ (self ):
272+ self .base_calls = 0
273+ self .derived_calls = 0
274+
275+ @wrapt .lru_cache
276+ def compute (self , x ):
277+ self .base_calls += 1
278+ return x * 2
279+
280+
281+ class OverrideDerived (OverrideBase ):
282+ @wrapt .lru_cache
283+ def compute (self , x ):
284+ self .derived_calls += 1
285+ return super ().compute (x ) + 1
286+
287+
288+ class TestOverriddenMethodWithSuper (unittest .TestCase ):
289+ def test_super_call_returns_correct_result (self ):
290+ obj = OverrideDerived ()
291+ self .assertEqual (obj .compute (10 ), 21 )
292+
293+ def test_super_call_does_not_recurse (self ):
294+ # Regression test: a subclass method decorated with lru_cache that
295+ # called the base method via super() used to recurse forever because
296+ # the base and derived methods shared a single per-instance cache
297+ # slot derived from the method name alone.
298+ obj = OverrideDerived ()
299+ try :
300+ result = obj .compute (10 )
301+ except RecursionError :
302+ self .fail ("super() call recursed instead of reaching base method" )
303+ self .assertEqual (result , 21 )
304+
305+ def test_both_bodies_execute_once (self ):
306+ obj = OverrideDerived ()
307+ obj .compute (10 )
308+ self .assertEqual (obj .derived_calls , 1 )
309+ self .assertEqual (obj .base_calls , 1 )
310+
311+ def test_base_and_derived_cached_independently (self ):
312+ obj = OverrideDerived ()
313+ obj .compute (10 )
314+ obj .compute (10 )
315+ # Second call is served from both caches, so neither body re-runs.
316+ self .assertEqual (obj .derived_calls , 1 )
317+ self .assertEqual (obj .base_calls , 1 )
318+ info = obj .compute .cache_info ()
319+ self .assertEqual (info .hits , 1 )
320+ self .assertEqual (info .misses , 1 )
321+
322+ def test_base_class_instance_unaffected (self ):
323+ obj = OverrideBase ()
324+ self .assertEqual (obj .compute (10 ), 20 )
325+ self .assertEqual (obj .base_calls , 1 )
326+ self .assertEqual (obj .derived_calls , 0 )
327+
328+ def test_separate_instances_have_separate_caches (self ):
329+ obj1 = OverrideDerived ()
330+ obj2 = OverrideDerived ()
331+ obj1 .compute (10 )
332+ obj2 .compute (10 )
333+ self .assertEqual (obj1 .derived_calls , 1 )
334+ self .assertEqual (obj2 .derived_calls , 1 )
335+ self .assertEqual (obj1 .base_calls , 1 )
336+ self .assertEqual (obj2 .base_calls , 1 )
337+
338+
270339class TestIntrospection (unittest .TestCase ):
271340 def test_function_name (self ):
272341 self .assertEqual (cached_function .__name__ , "cached_function" )
0 commit comments