Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions CHANGES.md
Original file line number Diff line number Diff line change
@@ -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))

Expand Down
34 changes: 34 additions & 0 deletions cloudpickle/cloudpickle.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
88 changes: 87 additions & 1 deletion tests/cloudpickle_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand Down
1 change: 1 addition & 0 deletions tests/cloudpickle_testpkg/_cloudpickle_testpkg/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,3 +46,4 @@ def g():

some_singleton = _SingletonClass()
T = typing.TypeVar("T")
MyInt = typing.NewType("MyInt", int)
1 change: 1 addition & 0 deletions tests/mock_local_folder/mod.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,3 +19,4 @@ def method(self):


LocalT = typing.TypeVar("LocalT")
LocalNewType = typing.NewType("LocalNewType", int)