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
5 changes: 5 additions & 0 deletions CHANGES.md
Original file line number Diff line number Diff line change
@@ -1,6 +1,11 @@
In development
==============

- Preserve explicit subclass attributes and methods that reference the same
objects as attributes on a base class, so updating the base after unpickling
does not change the subclass's overrides.
([issue#584](https://github.com/cloudpipe/cloudpickle/issues/584))

- Make pickling of functions depending on globals in notebook more
deterministic. ([PR#560](https://github.com/cloudpipe/cloudpickle/pull/560))

Expand Down
49 changes: 32 additions & 17 deletions cloudpickle/cloudpickle.py
Original file line number Diff line number Diff line change
Expand Up @@ -423,28 +423,30 @@ def _walk_global_ops(code):
yield instr.argval


class _BaseAttributeRef:
"""Reference an explicit class attribute to one of its direct bases."""

__slots__ = ("base_index",)

def __init__(self, base_index):
self.base_index = base_index


def _extract_class_dict(cls):
"""Retrieve a copy of the dict of a class without the inherited method."""
"""Copy a class's own attributes, including explicit overrides of its bases."""
# Hack to circumvent non-predictable memoization caused by string interning.
# See the inline comment in _class_setstate for details.
clsdict = {"".join(k): cls.__dict__[k] for k in sorted(cls.__dict__)}

if len(cls.__bases__) == 1:
inherited_dict = cls.__bases__[0].__dict__
else:
inherited_dict = {}
for base in reversed(cls.__bases__):
inherited_dict.update(base.__dict__)
to_remove = []
for name, value in clsdict.items():
try:
base_value = inherited_dict[name]
if value is base_value:
to_remove.append(name)
except KeyError:
pass
for name in to_remove:
clsdict.pop(name)
if name == "__module__":
continue
for index, base in enumerate(cls.__bases__):
if name in base.__dict__:
if value is base.__dict__[name]:
# Preserve the local binding without serializing the same
# value again, which might not be picklable by itself.
clsdict[name] = _BaseAttributeRef(index)
break
return clsdict


Expand Down Expand Up @@ -1179,8 +1181,21 @@ def _class_setstate(obj, state):
state, slotstate = state
registry = None
for attrname, attr in state.items():
if isinstance(attr, _BaseAttributeRef):
# A tracked dynamic class may already have its original binding.
if attrname in obj.__dict__:
continue
try:
attr = obj.__bases__[attr.base_index].__dict__[attrname]
except (IndexError, KeyError):
# The base may have changed since this class was pickled.
continue
if attrname == "_abc_impl":
registry = attr
elif attrname == "__module__" and obj.__dict__.get(attrname) == attr:
# Skeleton construction already sets __module__. Avoid invoking
# custom metaclass setters again unless the value has changed.
continue
else:
# Note: setting attribute names on a class automatically triggers their
# interning in CPython:
Expand Down
152 changes: 151 additions & 1 deletion tests/cloudpickle_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,17 +109,167 @@ def method_c(self):
return "c"

clsdict = _extract_class_dict(C)
expected_keys = ["C_CONSTANT", "__doc__", "method_c"]
expected_keys = ["C_CONSTANT", "__doc__", "__module__", "method_c"]
# New attribute in Python 3.13 beta 1
# https://github.com/python/cpython/pull/118475
if sys.version_info >= (3, 13):
expected_keys.insert(2, "__firstlineno__")
expected_keys.insert(4, "__static_attributes__")
assert list(clsdict.keys()) == expected_keys
assert clsdict["C_CONSTANT"] == 43
assert clsdict["__doc__"] is None
assert clsdict["method_c"](C()) == C().method_c()


@pytest.mark.parametrize("protocol", [2, cloudpickle.DEFAULT_PROTOCOL])
@pytest.mark.parametrize("multiple_inheritance", [False, True])
def test_class_explicit_overrides(protocol, multiple_inheritance):
# Start the worker before defining the classes so fork cannot copy their
# entries in cloudpickle's dynamic class tracker.
with subprocess_worker(protocol=protocol) as worker:

class Parent:
value = 1
inherited = 1

def method(self):
return "original"

class Mixin:
pass

bases = (Parent, Mixin) if multiple_inheritance else (Parent,)

class Child(*bases):
value = Parent.value
method = Parent.method

def check_overrides(child):
parent = child.__bases__[0]
parent.value = 2
parent.inherited = 2
parent.method = lambda self: "updated"
assert child.value == 1
assert child().method() == "original"
assert child.inherited == 2
assert "value" in child.__dict__
assert "method" in child.__dict__
assert "inherited" not in child.__dict__

worker.run(check_overrides, Child)
check_overrides(Child)


@pytest.mark.parametrize("protocol", [2, cloudpickle.DEFAULT_PROTOCOL])
@pytest.mark.parametrize("multiple_inheritance", [False, True])
def test_class_explicit_override_of_importable_unpicklable_attribute(
protocol, multiple_inheritance
):
testpkg = pytest.importorskip("_cloudpickle_testpkg")

with subprocess_worker(protocol=protocol) as worker:

class Mixin:
pass

bases = (
(Mixin, testpkg.BaseWithLock)
if multiple_inheritance
else (testpkg.BaseWithLock,)
)

class Child(*bases):
shared = testpkg.BaseWithLock.shared

def check_override(child):
assert "shared" in child.__dict__
base = next(base for base in child.__bases__ if "shared" in base.__dict__)
assert child.__dict__["shared"] is base.shared

worker.run(check_override, Child)
check_override(pickle_depickle(Child, protocol=protocol))


@pytest.mark.parametrize("protocol", [2, cloudpickle.DEFAULT_PROTOCOL])
def test_class_explicit_override_keeps_existing_binding(monkeypatch, protocol):
testpkg = pytest.importorskip("_cloudpickle_testpkg")

class Child(testpkg.BaseWithLock):
shared = testpkg.BaseWithLock.shared

original = Child.shared
payload = cloudpickle.dumps(Child, protocol=protocol)
monkeypatch.setattr(testpkg.BaseWithLock, "shared", object())

restored = pickle.loads(payload)
assert restored is Child
assert restored.__dict__["shared"] is original


@pytest.mark.parametrize("protocol", [2, cloudpickle.DEFAULT_PROTOCOL])
def test_class_explicit_override_resolves_changed_importable_base(protocol):
testpkg = pytest.importorskip("_cloudpickle_testpkg")

with subprocess_worker(protocol=protocol) as worker:

class Child(testpkg.BaseWithLock):
shared = testpkg.BaseWithLock.shared

def replace_base_attribute():
import _cloudpickle_testpkg

_cloudpickle_testpkg.BaseWithLock.shared = object()

def check_override(child):
assert "shared" in child.__dict__
assert child.__dict__["shared"] is child.__bases__[0].shared

worker.run(replace_base_attribute)
worker.run(check_override, Child)


@pytest.mark.parametrize("protocol", [2, cloudpickle.DEFAULT_PROTOCOL])
def test_class_module_set_during_construction(monkeypatch, protocol):
testpkg = pytest.importorskip("_cloudpickle_testpkg")

class ModuleOnceMeta(type):
__module__ = testpkg.__name__
__qualname__ = "ModuleOnceMeta"

def __setattr__(cls, name, value):
if name == "__module__" and cls.__dict__.get(name) == value:
raise TypeError("redundant module assignment")
super().__setattr__(name, value)

class Parent(metaclass=ModuleOnceMeta):
__module__ = testpkg.__name__
__qualname__ = "ModuleOnceParent"

monkeypatch.setattr(testpkg, "ModuleOnceMeta", ModuleOnceMeta, raising=False)
monkeypatch.setattr(testpkg, "ModuleOnceParent", Parent, raising=False)
assert pickle.loads(pickle.dumps(Parent)) is Parent

class Child(Parent):
__module__ = testpkg.__name__

restored = pickle_depickle(Child, protocol=protocol)
assert restored.__module__ == testpkg.__name__
assert restored.__bases__ == (Parent,)


@pytest.mark.parametrize("protocol", [2, cloudpickle.DEFAULT_PROTOCOL])
def test_class_module_restored_from_pickle(protocol):
class DynamicClass:
pass

original_module = DynamicClass.__module__
payload = cloudpickle.dumps(DynamicClass, protocol=protocol)
DynamicClass.__module__ = "changed_module"
restored = pickle.loads(payload)
assert restored is DynamicClass
assert restored.__module__ == original_module


class CloudPickleTest(unittest.TestCase):
protocol = cloudpickle.DEFAULT_PROTOCOL

Expand Down
5 changes: 5 additions & 0 deletions tests/cloudpickle_testpkg/_cloudpickle_testpkg/__init__.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import threading
import typing
from . import mod # noqa

Expand All @@ -10,6 +11,10 @@ def package_function():
global_variable = "some global variable"


class BaseWithLock:
shared = threading.Lock()


def package_function_with_global():
global global_variable
return global_variable
Expand Down