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
14 changes: 13 additions & 1 deletion cloudpickle/cloudpickle.py
Original file line number Diff line number Diff line change
Expand Up @@ -671,8 +671,20 @@ def _make_dict_items(obj, is_ordered=False):

def _class_getnewargs(obj):
type_kwargs = {}
bases = _get_bases(obj)
if "__module__" in obj.__dict__:
type_kwargs["__module__"] = obj.__module__
# Slots must be present when the class is created: assigning __slots__
# afterward cannot restore the instance layout. NamedTuple creates its own
# slots and rejects an explicit __slots__ entry in the class namespace.
named_tuple_bases = (
typing.NamedTuple,
getattr(sys.modules.get("typing_extensions"), "NamedTuple", None),
)
if "__slots__" in obj.__dict__ and not any(
base in named_tuple_bases for base in bases
):
type_kwargs["__slots__"] = obj.__slots__

__dict__ = obj.__dict__.get("__dict__", None)
if isinstance(__dict__, property):
Expand All @@ -681,7 +693,7 @@ def _class_getnewargs(obj):
return (
type(obj),
obj.__name__,
_get_bases(obj),
bases,
type_kwargs,
_get_or_create_tracker_id(obj),
None,
Expand Down
32 changes: 32 additions & 0 deletions tests/cloudpickle_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -2313,6 +2313,38 @@ def test_type_hint(self):
t = typing.Union[list, int]
assert pickle_depickle(t) == t

def test_typing_extensions_namedtuple(self):
typing_extensions = pytest.importorskip("typing_extensions")

class MyTuple(typing_extensions.NamedTuple):
value: int

restored = subprocess_pickle_echo(MyTuple(42), protocol=self.protocol)
assert restored.value == 42
assert not hasattr(restored, "__dict__")

def test_slotted_class_layout_in_subprocess(self):
class Base:
__slots__ = ("base_value",)

class Child(Base):
__slots__ = ("child_value",)

obj = Child()
obj.base_value = 1
obj.child_value = 2

def check_layout(restored):
assert restored.base_value == 1
assert restored.child_value == 2
assert not hasattr(restored, "__dict__")
with pytest.raises(AttributeError):
restored.extra = 3
return restored.base_value + restored.child_value

with subprocess_worker(protocol=self.protocol) as worker:
assert worker.run(check_layout, obj) == 3

def test_instance_with_slots(self):
for slots in [["registered_attribute"], "registered_attribute"]:

Expand Down