From 5823b29bbb88939ee6de9659ba0dcd425ec94ef3 Mon Sep 17 00:00:00 2001 From: Zhufeng Pan Date: Sat, 26 Sep 2026 01:24:46 -0700 Subject: [PATCH] Keep registered immutable objects alive to prevent object ID reuse. Store `id(value) -> value` in `_IMMUTABLE_OBJECT_IDS` and `_FUNCTIONS_WITH_IMMUTABLE_RETURN_VALUES` instead of storing only `id(value)` in sets, so registered objects are not garbage collected and their memory addresses cannot be reused by newly allocated objects. PiperOrigin-RevId: 988745384 --- fiddle/_src/daglish_extensions.py | 15 +++++++++------ fiddle/_src/daglish_extensions_test.py | 22 +++++++++++++++++++--- 2 files changed, 28 insertions(+), 9 deletions(-) diff --git a/fiddle/_src/daglish_extensions.py b/fiddle/_src/daglish_extensions.py index b0362de1..5e11fe48 100644 --- a/fiddle/_src/daglish_extensions.py +++ b/fiddle/_src/daglish_extensions.py @@ -26,23 +26,25 @@ from fiddle._src import config as config_lib from fiddle._src import daglish -_IMMUTABLE_OBJECT_IDS = set() +# Note: we include `value` in the dict to keep it alive, to ensure that +# `id(value)` remains valid and is not reused after garbage collection. +_IMMUTABLE_OBJECT_IDS = {} def register_immutable(value: Any) -> None: """Registers a certain type to be immutable.""" - _IMMUTABLE_OBJECT_IDS.add(id(value)) + _IMMUTABLE_OBJECT_IDS[id(value)] = value -# Similar set of IDs of functions/types with immutable return values (for this +# Similar dict of IDs of functions/types with immutable return values (for this # case, maybe it's OK to replace IDs with just the functions/types?). -_FUNCTIONS_WITH_IMMUTABLE_RETURN_VALUES = set() +_FUNCTIONS_WITH_IMMUTABLE_RETURN_VALUES = {} def register_function_with_immutable_return_value( - fn_or_cls: Union[Type[Any], Callable[..., Any]] + fn_or_cls: Union[Type[Any], Callable[..., Any]], ) -> None: - _FUNCTIONS_WITH_IMMUTABLE_RETURN_VALUES.add(id(fn_or_cls)) + _FUNCTIONS_WITH_IMMUTABLE_RETURN_VALUES[id(fn_or_cls)] = fn_or_cls # TODO(b/285146396): Register the tensorflow dtypes as immutable. @@ -117,6 +119,7 @@ def is_unshareable(value: Any) -> bool: ) ) + _PATH_PART = re.compile( "(?:{})".format( "|".join([ diff --git a/fiddle/_src/daglish_extensions_test.py b/fiddle/_src/daglish_extensions_test.py index 32509f16..069bb643 100644 --- a/fiddle/_src/daglish_extensions_test.py +++ b/fiddle/_src/daglish_extensions_test.py @@ -15,7 +15,6 @@ """Tests for daglish_extensions.""" - import dataclasses from typing import Any, NamedTuple @@ -51,13 +50,24 @@ class DaglishExtensionsTest(parameterized.TestCase): def test_register_immutable(self): obj = object() self.assertFalse(daglish_extensions.is_immutable(obj)) + self.assertFalse(daglish_extensions.is_unshareable(obj)) daglish_extensions.register_immutable(obj) self.assertTrue(daglish_extensions.is_immutable(obj)) + self.assertTrue(daglish_extensions.is_unshareable(obj)) + del obj + new_obj = object() + self.assertFalse(daglish_extensions.is_immutable(new_obj)) + self.assertFalse(daglish_extensions.is_unshareable(new_obj)) def test_register_function_with_immutable_return_value(self): - def fn(x: int) -> int: - return x + def make_fn(): + def fn(x): + return x + + return fn + + fn = make_fn() config = fdl.Config(fn, 3) self.assertFalse(daglish_extensions.is_immutable(config)) self.assertFalse(daglish_extensions.is_unshareable(config)) @@ -69,6 +79,12 @@ def fn(x: int) -> int: self.assertFalse(daglish_extensions.is_immutable(config)) self.assertTrue(daglish_extensions.is_unshareable(config)) + fn2 = make_fn() + daglish_extensions.register_function_with_immutable_return_value(fn2) + del fn2 + fn3 = make_fn() + self.assertFalse(daglish_extensions.is_unshareable(fdl.Config(fn3, 3))) + @parameterized.parameters( { "path": ".foo.bar",