From 949660323bbd7c2c412e9571968700db7f719a82 Mon Sep 17 00:00:00 2001 From: Charan Rathore <180254320+charan-rathore@users.noreply.github.com> Date: Thu, 1 Oct 2026 13:33:02 +0530 Subject: [PATCH] Fix pickling of NewType instances defined in __main__ on Python 3.10+ --- CHANGES.md | 4 + cloudpickle/cloudpickle.py | 34 +++++++ tests/cloudpickle_test.py | 88 ++++++++++++++++++- .../_cloudpickle_testpkg/__init__.py | 1 + tests/mock_local_folder/mod.py | 1 + 5 files changed, 127 insertions(+), 1 deletion(-) diff --git a/CHANGES.md b/CHANGES.md index a6b0b443..b1878b0a 100644 --- a/CHANGES.md +++ b/CHANGES.md @@ -1,6 +1,10 @@ In development ============== +- Fix pickling of `typing.NewType` instances defined in the `__main__` + module on Python 3.10 and newer. + ([issue #520](https://github.com/cloudpipe/cloudpickle/issues/520)) + - Make pickling of functions depending on globals in notebook more deterministic. ([PR#560](https://github.com/cloudpipe/cloudpickle/pull/560)) diff --git a/cloudpickle/cloudpickle.py b/cloudpickle/cloudpickle.py index 08882306..773df53b 100644 --- a/cloudpickle/cloudpickle.py +++ b/cloudpickle/cloudpickle.py @@ -630,6 +630,36 @@ def _typevar_reduce(obj): return (getattr, module_and_name) +def _make_newtype(name, qualname, module, supertype, class_tracker_id): + nt = typing.NewType(name, supertype) + nt.__qualname__ = qualname + nt.__module__ = module + return _lookup_class_or_track(class_tracker_id, nt) + + +def _decompose_newtype(obj): + return ( + obj.__name__, + obj.__qualname__, + obj.__module__, + obj.__supertype__, + _get_or_create_tracker_id(obj), + ) + + +def _newtype_reduce(obj): + # NewType instances require the module information hence why we + # are not using the _should_pickle_by_reference directly + module_and_name = _lookup_module_and_qualname(obj, name=obj.__qualname__) + + if module_and_name is None: + return (_make_newtype, _decompose_newtype(obj)) + elif _is_registered_pickle_by_value(module_and_name[0]): + return (_make_newtype, _decompose_newtype(obj)) + + return (getattr, module_and_name) + + def _get_bases(typ): if "__orig_bases__" in getattr(typ, "__dict__", {}): # For generic types (see PEP 560) @@ -1253,6 +1283,10 @@ class Pickler(pickle.Pickler): _dispatch_table[types.MappingProxyType] = _mappingproxy_reduce _dispatch_table[weakref.WeakSet] = _weakset_reduce _dispatch_table[typing.TypeVar] = _typevar_reduce + if isinstance(typing.NewType, type): + # NewType was a function before Python 3.10: its instances only + # exist as instances of typing.NewType from Python 3.10 onwards. + _dispatch_table[typing.NewType] = _newtype_reduce _dispatch_table[_collections_abc.dict_keys] = _dict_keys_reduce _dispatch_table[_collections_abc.dict_values] = _dict_values_reduce _dispatch_table[_collections_abc.dict_items] = _dict_items_reduce diff --git a/tests/cloudpickle_test.py b/tests/cloudpickle_test.py index e2097d1c..993a5f91 100644 --- a/tests/cloudpickle_test.py +++ b/tests/cloudpickle_test.py @@ -2591,6 +2591,85 @@ def test_pickle_importable_typevar(self): assert AnyStr is pickle_depickle(AnyStr, protocol=self.protocol) + newtype_only = pytest.mark.skipif( + sys.version_info < (3, 10), + reason="typing.NewType is a class from Python 3.10 onwards only", + ) + + @newtype_only + def test_pickle_dynamic_newtype(self): + MyInt = typing.NewType("MyInt", int) + depickled_myint = pickle_depickle(MyInt, protocol=self.protocol) + assert depickled_myint is MyInt + + @newtype_only + def test_pickle_dynamic_newtype_tracking(self): + MyInt = typing.NewType("MyInt", int) + MyInt2 = subprocess_pickle_echo(MyInt, protocol=self.protocol) + assert MyInt is MyInt2 + + @newtype_only + def test_pickle_dynamic_newtype_memoization(self): + MyInt = typing.NewType("MyInt", int) + depickled_1, depickled_2 = pickle_depickle( + (MyInt, MyInt), protocol=self.protocol + ) + assert depickled_1 is depickled_2 + + @newtype_only + def test_pickle_newtype_with_newtype_supertype(self): + MyInt = typing.NewType("MyInt", int) + MyOtherInt = typing.NewType("MyOtherInt", MyInt) + MyInt2, MyOtherInt2 = subprocess_pickle_echo( + (MyInt, MyOtherInt), protocol=self.protocol + ) + assert MyInt is MyInt2 + assert MyOtherInt is MyOtherInt2 + assert MyOtherInt2.__supertype__ is MyInt2 + + @newtype_only + def test_pickle_importable_newtype(self): + _cloudpickle_testpkg = pytest.importorskip("_cloudpickle_testpkg") + MyInt = pickle_depickle(_cloudpickle_testpkg.MyInt, protocol=self.protocol) + assert MyInt is _cloudpickle_testpkg.MyInt + + @newtype_only + def test_pickle_newtype_in_main_module(self): + # Non-regression test for + # https://github.com/cloudpipe/cloudpickle/issues/520 + with tempfile.TemporaryDirectory() as tmpdir: + pickle_file = os.path.join(tmpdir, "newtype.pickle") + assert_run_python_script( + textwrap.dedent( + """ + import cloudpickle + from typing import NewType + + RunDict = NewType("RunDict", dict) + with open(%r, "wb") as f: + f.write(cloudpickle.dumps(RunDict, protocol=%d)) + """ + % (pickle_file, self.protocol) + ) + ) + assert_run_python_script( + textwrap.dedent( + """ + import pickle + + with open(%r, "rb") as f: + RunDict = pickle.load(f) + + assert RunDict.__name__ == "RunDict" + assert RunDict.__qualname__ == "RunDict" + assert RunDict.__module__ == "__main__" + assert RunDict.__supertype__ is dict + assert RunDict({"a": 1}) == {"a": 1} + """ + % pickle_file + ) + ) + def test_generic_type(self): T = typing.TypeVar("T") @@ -2791,7 +2870,12 @@ def test_pickle_constructs_from_module_registered_for_pickling_by_value( # The constructs whose pickling mechanism is changed using # register_pickle_by_value are functions, classes, TypeVar and # modules. - from mock_local_folder.mod import local_function, LocalT, LocalClass + from mock_local_folder.mod import ( + local_function, + LocalT, + LocalClass, + LocalNewType, + ) # Make sure the module/constructs are unimportable in the # worker. @@ -2809,6 +2893,8 @@ def test_pickle_constructs_from_module_registered_for_pickling_by_value( assert w.run(lambda: local_function()) == local_function() # typevar assert w.run(lambda: LocalT.__name__) == LocalT.__name__ + # newtype + assert w.run(lambda: LocalNewType.__name__) == LocalNewType.__name__ # classes assert w.run(lambda: LocalClass().method()) == LocalClass().method() # modules diff --git a/tests/cloudpickle_testpkg/_cloudpickle_testpkg/__init__.py b/tests/cloudpickle_testpkg/_cloudpickle_testpkg/__init__.py index 2051e4c0..4f74e49f 100644 --- a/tests/cloudpickle_testpkg/_cloudpickle_testpkg/__init__.py +++ b/tests/cloudpickle_testpkg/_cloudpickle_testpkg/__init__.py @@ -46,3 +46,4 @@ def g(): some_singleton = _SingletonClass() T = typing.TypeVar("T") +MyInt = typing.NewType("MyInt", int) diff --git a/tests/mock_local_folder/mod.py b/tests/mock_local_folder/mod.py index 517d5013..3778d23f 100644 --- a/tests/mock_local_folder/mod.py +++ b/tests/mock_local_folder/mod.py @@ -19,3 +19,4 @@ def method(self): LocalT = typing.TypeVar("LocalT") +LocalNewType = typing.NewType("LocalNewType", int)