diff --git a/docs/changes.rst b/docs/changes.rst index cbc50126..4db24835 100644 --- a/docs/changes.rst +++ b/docs/changes.rst @@ -26,17 +26,33 @@ v0.17 ``Series`` or ``DataFrame`` using one, could not be saved. Such objects are now saved through ``__getstate__``/``__dict__`` again, as before v0.12.0. :pr:`550` by `Adrin Jalali`_. -- Objects which contain a reference to themselves, directly or through - their attributes, dicts, lists, or sets, can now be saved and loaded. They - used to fail with a ``RecursionError``; this affected for instance the - discrete distributions of ``scipy.stats`` and fitted - :class:`sklearn.cluster.Birch` models, which is no longer listed as - unsupported. Such a reference through any other type, e.g. a tuple, raises - an ``UnsupportedTypeException`` when saving. Bound methods inherited from a - class defined in another module can now be loaded, they used to be - rejected as corrupted. A file in which nodes of different types share an - ``__id__`` is now rejected when loading instead of silently loading one of - them in place of the other. :pr:`549` by `Adrin Jalali`_. +- Objects which contain a reference to themselves can now be saved and + loaded, and subclasses of the built-in containers are persisted the way + pickle does it. :pr:`549` and :pr:`554` by `Adrin Jalali`_. + + - Circular references: objects which refer to themselves, directly or + through their attributes, dicts, lists, or sets, including subclasses of + these, can now be saved and loaded. They used to fail with a + ``RecursionError``; this affected for instance the discrete distributions + of ``scipy.stats`` and fitted :class:`sklearn.cluster.Birch` models, which + is no longer listed as unsupported. Such a reference through any other + type, e.g. a tuple, raises an ``UnsupportedTypeException`` when saving. + + - Container subclasses: subclasses of ``dict``, ``list`` and ``set`` are now + saved and loaded the way pickle does it for dicts and lists. The instance + is created with ``__new__`` and filled in place instead of through its + constructor, so a subclass whose constructor requires arguments can now + be loaded, and its instance attributes, which used to be dropped, are + saved and restored. Subclasses of ``set`` and of + ``collections.defaultdict`` are treated the same way, although pickle + builds them through their constructor, so that items which refer back to + the container can be loaded. Files written with an earlier protocol load + as before. + + - Loading: bound methods inherited from a class defined in another module + can now be loaded, they used to be rejected as corrupted. A file in which + nodes of different types share an ``__id__`` is now rejected instead of + silently loading one of them in place of the other. - Restore the ``skops`` command line entry point. It was declared in ``setup.py`` and lost when the packaging moved to ``pyproject.toml`` in v0.11.0, so ``skops convert`` and ``skops update`` had not been available @@ -48,6 +64,12 @@ v0.17 could not be loaded. The file format now stores the keyword arguments and the persistence protocol is bumped to 3; files written with an earlier protocol load as before. :pr:`551` by `Adrin Jalali`_. +- Fix loading of ``collections.defaultdict`` objects whose keys are not + strings, which failed with a ``TypeError``, and of objects whose class + defines ``__slots__``, which failed with an ``AttributeError`` on Python + 3.11 and later. Subclasses of ``defaultdict`` are now loaded as their own + type, with their instance attributes; they used to be loaded as a plain + ``defaultdict``. :pr:`554` by `Adrin Jalali`_. v0.16 ----- diff --git a/skops/io/_audit.py b/skops/io/_audit.py index a3598eed..fb1240d8 100644 --- a/skops/io/_audit.py +++ b/skops/io/_audit.py @@ -173,8 +173,8 @@ def construct(self) -> Any: A node is reached again while its own ``_construct`` runs when the saved object contained a reference to itself, directly or through its - children. ``ObjectNode``, ``DictNode``, ``ListNode`` and ``SetNode`` - support this by storing the instance in ``_constructed`` before + children. ``ObjectNode``, ``DictNode``, ``DefaultDictNode``, ``ListNode`` + and ``SetNode`` support this by storing the instance in ``_constructed`` before constructing the children, so the second call returns the partially constructed instance, like pickle does. Any other node raises instead of recursing until the interpreter gives up. diff --git a/skops/io/_general.py b/skops/io/_general.py index 7e8c2136..5eb65af2 100644 --- a/skops/io/_general.py +++ b/skops/io/_general.py @@ -40,6 +40,73 @@ arepr.maxstring = 24 +def _get_attrs_state(obj: Any, save_context: SaveContext) -> dict[str, Any] | None: + """Return the state of the instance attributes of a container, if any. + + A plain dict, list or set has no instance attributes. An instance of a + subclass has a ``__dict__``, unless the class defines ``__slots__``, and + its attributes are read the way ``object_get_state`` reads them: through + ``__getstate__`` when the object has one, and from ``__dict__`` otherwise. + Every object has a ``__getstate__`` since Python 3.11, and the default one + returns ``None`` when there is nothing to save. A class which defines its + own decides what is worth saving, an empty dict included, and gets it back + through ``__setstate__``, as with pickle. ``None`` is returned when there + is nothing to save, so that the state of an object without attributes is + the same on every Python version. + """ + if hasattr(obj, "__getstate__"): + attrs = obj.__getstate__() + else: + # Python < 3.11, where only a custom ``__getstate__`` exists: an empty + # ``__dict__`` is what the default ``__getstate__`` reports as ``None``. + attrs = getattr(obj, "__dict__", None) or None + if attrs is None: + return None + return get_state(attrs, save_context) + + +def _get_attrs_tree( + state: dict[str, Any], load_context: LoadContext, trusted: TrustedTypes | None +) -> Node | None: + """Return the node of the ``attrs`` entry of ``state``, if there is one. + + The entry is written by ``_get_attrs_state``, for instances which have + attributes. + """ + attrs = state.get("attrs") + if attrs is None: + return None + return get_tree(attrs, load_context, trusted=trusted) + + +def _set_attrs(instance: Any, attrs: Node | None) -> None: + """Give ``instance`` its attributes back from the node of their state. + + ``ObjectNode``, ``DictNode``, ``ListNode`` and ``SetNode`` create the + instance with ``__new__`` and then restore its attributes the way pickle + does: through ``__setstate__`` when the instance has one, and otherwise + by updating its ``__dict__``. The default ``__getstate__`` of an object + whose class defines ``__slots__`` returns a ``(dict_state, slots_state)`` + tuple instead of a dict, and the slot values are set one by one. + """ + if attrs is None: + return + state = attrs.construct() + if state is None: + return + if hasattr(instance, "__setstate__"): + instance.__setstate__(state) + return + slots_state = None + if isinstance(state, tuple) and len(state) == 2: + state, slots_state = state + if state: + instance.__dict__.update(state) + if slots_state: + for name, value in slots_state.items(): + setattr(instance, name, value) + + def dict_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]: res = { "__class__": obj.__class__.__name__, @@ -58,6 +125,9 @@ def dict_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]: content[key] = get_state(value, save_context) res["content"] = content res["key_types"] = key_types + attrs = _get_attrs_state(obj, save_context) + if attrs is not None: + res["attrs"] = attrs return res @@ -78,16 +148,25 @@ def __init__( for key, value in state["content"].items() } self.children = {"key_types": self.key_types, "content": self.content} + # the instance attributes of a subclass, see ``_get_attrs_state`` + self.attrs = _get_attrs_tree(state, load_context, trusted) + if self.attrs is not None: + self.children["attrs"] = self.attrs def _construct(self): - content = gettype(self.module_name, self.class_name)() - # Make the dict available to children which refer back to it, see + cls: type[object] = gettype(self.module_name, self.class_name) + # As pickle does, the instance is created with ``__new__`` instead of + # through its constructor, filled in place and then given its + # attributes back. It is stored before its items are constructed, so + # that items which refer back to it get the instance, see # ``Node.construct``. - self._constructed = content + instance: Any = cls.__new__(cls) + self._constructed = instance key_types = self.key_types.construct() for k_type, (key, val) in zip(key_types, self.content.items()): - content[k_type(key)] = val.construct() - return content + instance[k_type(key)] = val.construct() + _set_attrs(instance, self.attrs) + return instance def defaultdict_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]: @@ -102,6 +181,9 @@ def defaultdict_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any] content["main"] = get_state(dict(obj), save_context) content["default_factory"] = get_state(obj.default_factory, save_context) res["content"] = content + attrs = _get_attrs_state(obj, save_context) + if attrs is not None: + res["attrs"] = attrs return res @@ -113,7 +195,7 @@ def __init__( trusted: TrustedTypes | None = None, ) -> None: super().__init__(state, load_context, trusted) - self.trusted = ["collections.defaultdict"] + self.trusted = self._get_trusted(trusted, ["collections.defaultdict"]) self.main = get_tree( state["content"]["main"], load_context, @@ -124,10 +206,24 @@ def __init__( state["content"]["default_factory"], load_context, trusted=trusted ) self.children = {"main": self.main, "default_factory": self.default_factory} + # the instance attributes of a subclass, see ``_get_attrs_state`` + self.attrs = _get_attrs_tree(state, load_context, trusted) + if self.attrs is not None: + self.children["attrs"] = self.attrs def _construct(self): - instance = defaultdict(**self.main.construct()) + cls: type[object] = gettype(self.module_name, self.class_name) + # Like ``DictNode``: the instance is created with ``__new__``, stored, + # given its default factory, filled in place and then given its + # attributes back. Pickle builds a defaultdict through its constructor + # with the default factory instead, which a subclass whose constructor + # takes other arguments does not accept. + instance: Any = cls.__new__(cls) + self._constructed = instance instance.default_factory = self.default_factory.construct() + for key, value in self.main.construct().items(): + instance[key] = value + _set_attrs(instance, self.attrs) return instance @@ -140,6 +236,9 @@ def list_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]: content = [get_state(value, save_context) for value in obj] res["content"] = content + attrs = _get_attrs_state(obj, save_context) + if attrs is not None: + res["attrs"] = attrs return res @@ -156,22 +255,23 @@ def __init__( get_tree(value, load_context, trusted=trusted) for value in state["content"] ] self.children = {"content": self.content} + # the instance attributes of a subclass, see ``_get_attrs_state`` + self.attrs = _get_attrs_tree(state, load_context, trusted) + if self.attrs is not None: + self.children["attrs"] = self.attrs def _construct(self): - content_type = gettype(self.module_name, self.class_name) - if content_type is not list: - # Subclasses are built from their items through their own - # constructor, as before, so there is no instance to hand out - # before it is complete: a reference back to a list subclass is not - # supported, and ``get_state`` refuses it when saving. - return content_type([item.construct() for item in self.content]) - - # Fill a plain list in place and make it available to children which - # refer back to it, see ``Node.construct``. - content: list[Any] = [] - self._constructed = content - content.extend(item.construct() for item in self.content) - return content + cls: type[object] = gettype(self.module_name, self.class_name) + # As pickle does, the instance is created with ``__new__`` instead of + # through its constructor, filled in place and then given its + # attributes back. It is stored before its items are constructed, so + # that items which refer back to it get the instance, see + # ``Node.construct``. + instance: Any = cls.__new__(cls) + self._constructed = instance + instance.extend([item.construct() for item in self.content]) + _set_attrs(instance, self.attrs) + return instance def set_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]: @@ -182,6 +282,9 @@ def set_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]: } content = [get_state(value, save_context) for value in obj] res["content"] = content + attrs = _get_attrs_state(obj, save_context) + if attrs is not None: + res["attrs"] = attrs return res @@ -198,22 +301,24 @@ def __init__( get_tree(value, load_context, trusted=trusted) for value in state["content"] ] self.children = {"content": self.content} + # the instance attributes of a subclass, see ``_get_attrs_state`` + self.attrs = _get_attrs_tree(state, load_context, trusted) + if self.attrs is not None: + self.children["attrs"] = self.attrs def _construct(self): - content_type = gettype(self.module_name, self.class_name) - if content_type is not set: - # Subclasses are built from their items through their own - # constructor, as before, so there is no instance to hand out - # before it is complete: a reference back to a set subclass is not - # supported, and ``get_state`` refuses it when saving. - return content_type([item.construct() for item in self.content]) - - # Fill a plain set in place and make it available to children which - # refer back to it, see ``Node.construct``. - content: set[Any] = set() - self._constructed = content - content.update(item.construct() for item in self.content) - return content + cls: type[object] = gettype(self.module_name, self.class_name) + # Like ``ListNode``: the instance is created with ``__new__``, stored, + # filled in place and then given its attributes back. Pickle builds a + # set subclass through its constructor instead, which would leave no + # instance to hand out to items which refer back to it. The items are + # added with the built-in ``set.update``, as the constructor would, so + # that an overridden ``update`` does not run on a bare instance. + instance: Any = cls.__new__(cls) + self._constructed = instance + set.update(instance, [item.construct() for item in self.content]) + _set_attrs(instance, self.attrs) + return instance def tuple_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]: @@ -578,19 +683,7 @@ def _construct(self): # Make the instance available to attributes which refer back to it, see # ``Node.construct``. self._constructed = instance - - attrs_node = self.attrs - if attrs_node is None: - # nothing more to do - return instance - - attrs = attrs_node.construct() - if attrs is not None: - if hasattr(instance, "__setstate__"): - instance.__setstate__(attrs) - else: - instance.__dict__.update(attrs) - + _set_attrs(instance, self.attrs) return instance diff --git a/skops/io/_utils.py b/skops/io/_utils.py index 699bb473..624454b1 100644 --- a/skops/io/_utils.py +++ b/skops/io/_utils.py @@ -227,16 +227,19 @@ def _get_state(obj, save_context: SaveContext): raise TypeError(f"Getting the state of type {type(obj)} is not supported yet") -def _supports_circular_reference(value: Any, state: dict[str, Any]) -> bool: - """Whether a reference to ``value`` from inside its own state can be loaded. +def _supports_circular_reference(state: dict[str, Any]) -> bool: + """Whether a reference to the object of ``state`` from inside it can be loaded. This mirrors the ``Node`` classes whose ``_construct`` registers the instance before constructing its children, see ``Node.construct``. - ``ListNode`` and ``SetNode`` only do so for plain lists and sets. """ - if type(value) in (list, set): - return True - return state["__loader__"] in ("DictNode", "ObjectNode") + return state["__loader__"] in ( + "DictNode", + "DefaultDictNode", + "ListNode", + "SetNode", + "ObjectNode", + ) def get_state(value, save_context: SaveContext) -> dict[str, Any]: @@ -264,9 +267,7 @@ def get_state(value, save_context: SaveContext) -> dict[str, Any]: save_context.in_progress[__id__] = False try: res = _get_state(value, save_context) - if save_context.in_progress[__id__] and not _supports_circular_reference( - value, res - ): + if save_context.in_progress[__id__] and not _supports_circular_reference(res): raise UnsupportedTypeException( f"Objects of type {type(value).__name__} which contain a" " reference to themselves are not supported yet." diff --git a/skops/io/old/_general_v2.py b/skops/io/old/_general_v2.py index b2256041..fa6a324c 100644 --- a/skops/io/old/_general_v2.py +++ b/skops/io/old/_general_v2.py @@ -3,8 +3,9 @@ import operator from typing import Any +from skops.io import _general from skops.io._audit import Node, get_tree -from skops.io._utils import LoadContext, TrustedTypes +from skops.io._utils import LoadContext, TrustedTypes, gettype PROTOCOL = 2 @@ -38,7 +39,112 @@ def _construct(self): return op(*attrs) -# tuples of type and function that creates the instance of that type +# Protocol 2 readers for dicts, lists and sets. Protocol 3 added an ``attrs`` +# entry to their state, holding the instance attributes of a subclass, and the +# current nodes in ``skops.io._general`` create every instance with +# ``__new__``, fill it in place and restore its attributes, as pickle does for +# dicts and lists; this also lets an item refer back to a list or set subclass +# which holds it. Files up to protocol 2 have no ``attrs`` entry and are read +# here as before: a subclass is built through its constructor, a dict subclass +# with no arguments and then filled, a list or set subclass from its items. +# The classes derive from the current nodes only to pass the ``allowed_types`` +# check of the nodes which require a ``DictNode``, ``ListNode`` or ``SetNode`` +# child, e.g. the key types of a ``DictNode`` or the keyword arguments of a +# ``PartialNode``; they override everything they inherit. +class ListNode(_general.ListNode): + def __init__( + self, + state: dict[str, Any], + load_context: LoadContext, + trusted: TrustedTypes | None = None, + ) -> None: + Node.__init__(self, state, load_context, trusted) + self.trusted = self._get_trusted(trusted, [list]) + self.content = [ + get_tree(value, load_context, trusted=trusted) for value in state["content"] + ] + self.children = {"content": self.content} + + def _construct(self): + content_type = gettype(self.module_name, self.class_name) + if content_type is not list: + return content_type([item.construct() for item in self.content]) + + # Fill a plain list in place and make it available to children which + # refer back to it, see ``Node.construct``. + content: list[Any] = [] + self._constructed = content + content.extend(item.construct() for item in self.content) + return content + + +class SetNode(_general.SetNode): + def __init__( + self, + state: dict[str, Any], + load_context: LoadContext, + trusted: TrustedTypes | None = None, + ) -> None: + Node.__init__(self, state, load_context, trusted) + self.trusted = self._get_trusted(trusted, [set]) + self.content = [ + get_tree(value, load_context, trusted=trusted) for value in state["content"] + ] + self.children = {"content": self.content} + + def _construct(self): + content_type = gettype(self.module_name, self.class_name) + if content_type is not set: + return content_type([item.construct() for item in self.content]) + + # Fill a plain set in place and make it available to children which + # refer back to it, see ``Node.construct``. + content: set[Any] = set() + self._constructed = content + content.update(item.construct() for item in self.content) + return content + + +class DictNode(_general.DictNode): + def __init__( + self, + state: dict[str, Any], + load_context: LoadContext, + trusted: TrustedTypes | None = None, + ) -> None: + Node.__init__(self, state, load_context, trusted) + self.trusted = self._get_trusted(trusted, [dict, "collections.OrderedDict"]) + self.key_types = get_tree( + state["key_types"], load_context, trusted=trusted, allowed_types=(ListNode,) + ) + self.content = { + key: get_tree(value, load_context, trusted=trusted) + for key, value in state["content"].items() + } + self.children = {"key_types": self.key_types, "content": self.content} + + def _construct(self): + content = gettype(self.module_name, self.class_name)() + # Make the dict available to children which refer back to it, see + # ``Node.construct``. + self._constructed = content + key_types = self.key_types.construct() + for k_type, (key, val) in zip(key_types, self.content.items()): + content[k_type(key)] = val.construct() + return content + + +# ``get_tree`` looks a node up by the protocol of the file, and falls back to +# the current node when nothing is registered for that protocol. The state of +# the objects above did not change between protocol 0 and 2, so their readers +# are registered for every one of these protocols. NODE_TYPE_MAPPING = { - ("OperatorFuncNode", PROTOCOL): OperatorFuncNode, + (loader, protocol): node + for loader, node in ( + ("OperatorFuncNode", OperatorFuncNode), + ("DictNode", DictNode), + ("ListNode", ListNode), + ("SetNode", SetNode), + ) + for protocol in range(PROTOCOL + 1) } diff --git a/skops/io/tests/test_persist.py b/skops/io/tests/test_persist.py index 8debe27a..6df4be7d 100644 --- a/skops/io/tests/test_persist.py +++ b/skops/io/tests/test_persist.py @@ -53,7 +53,7 @@ StandardScaler, ) from sklearn.tree import DecisionTreeClassifier -from sklearn.utils import all_estimators, check_random_state +from sklearn.utils import Bunch, all_estimators, check_random_state from sklearn.utils._testing import SkipTest, set_random_state from sklearn.utils.estimator_checks import ( _enforce_estimator_tags_X, @@ -65,6 +65,7 @@ import skops from skops.io import dump, dumps, get_untrusted_types, load, loads, visualize from skops.io._audit import NODE_TYPE_MAPPING, Node, get_tree +from skops.io._general import dict_get_state, list_get_state, set_get_state from skops.io._protocol import PROTOCOL from skops.io._sklearn import UNSUPPORTED_TYPES, loss_get_state from skops.io._trusted_types import ( @@ -75,7 +76,7 @@ SCIPY_UFUNC_TYPE_NAMES, SKLEARN_ESTIMATOR_TYPE_NAMES, ) -from skops.io._utils import LoadContext, _get_state, get_state, gettype +from skops.io._utils import LoadContext, _get_state, get_module, get_state, gettype from skops.io.exceptions import UnsupportedTypeException, UntrustedTypesFoundException from skops.io.tests._utils import ( assert_method_outputs_equal, @@ -1660,6 +1661,10 @@ def test_circular_reference_in_list(): assert loaded[1] is loaded +class DictSubclass(dict): + pass + + class ListSubclass(list): pass @@ -1668,23 +1673,252 @@ class SetSubclass(set): pass -@pytest.mark.parametrize("container_type", [ListSubclass, SetSubclass]) -def test_list_and_set_subclasses_round_trip(container_type): - # Only plain lists and sets are filled in place, to resolve references back - # to them. Their subclasses are constructed from the items, as before. - obj = container_type([1, 2, 3]) +class DictWithAttrs(dict): + """Dict subclass with a required constructor argument and an attribute.""" + + def __init__(self, name, items): + super().__init__(items) + self.name = name + + +class ListWithAttrs(list): + """List subclass with a required constructor argument and an attribute.""" + + def __init__(self, name, items): + super().__init__(items) + self.name = name + + +class SetWithAttrs(set): + """Set subclass with a required constructor argument and an attribute.""" + + def __init__(self, name, items): + super().__init__(items) + self.name = name + + +CONTAINER_SUBCLASS_CASES = [ + pytest.param(DictSubclass, DictWithAttrs, {"a": 1, "b": 2}, id="dict"), + pytest.param(ListSubclass, ListWithAttrs, [1, 2, 3], id="list"), + pytest.param(SetSubclass, SetWithAttrs, {1, 2, 3}, id="set"), +] + + +@pytest.mark.parametrize("plain_subclass, with_attrs, items", CONTAINER_SUBCLASS_CASES) +def test_container_subclasses_round_trip(plain_subclass, with_attrs, items): + obj = plain_subclass(items) + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded) is plain_subclass + assert loaded == obj + + +@pytest.mark.parametrize("plain_subclass, with_attrs, items", CONTAINER_SUBCLASS_CASES) +def test_container_subclasses_keep_attributes(plain_subclass, with_attrs, items): + # Like pickle, the instance of a subclass is created with __new__ and + # filled in place, so its constructor does not have to accept the items, + # and its attributes are saved and restored. + obj = with_attrs("foo", items) + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded) is with_attrs + assert loaded == obj + assert loaded.name == "foo" + + +def test_container_attrs_saved_only_when_present(): + # A plain dict, list or set has no instance attributes, and neither has an + # instance of a subclass without any, so their state has no "attrs" entry + # and is the same as before. The entry holds the attributes as a dict. + save_context = make_save_context() + assert "attrs" not in dict_get_state({"a": 1}, save_context) + assert "attrs" not in list_get_state([1], save_context) + assert "attrs" not in set_get_state({1}, save_context) + assert "attrs" not in dict_get_state(DictSubclass({"a": 1}), save_context) + assert "attrs" not in list_get_state(ListSubclass([1]), save_context) + assert "attrs" not in set_get_state(SetSubclass([1]), save_context) + state = dict_get_state(DictWithAttrs("foo", {"a": 1}), save_context) + assert state["attrs"]["__loader__"] == "DictNode" + state = list_get_state(ListWithAttrs("foo", [1]), save_context) + assert state["attrs"]["__loader__"] == "DictNode" + state = set_get_state(SetWithAttrs("foo", [1]), save_context) + assert state["attrs"]["__loader__"] == "DictNode" + + +class ListWithCustomState(list): + """List subclass whose empty custom state must still reach __setstate__.""" + + def __getstate__(self): + return {} + + def __setstate__(self, state): + self.restored = state + + +def test_container_empty_custom_state_is_saved(): + # A custom __getstate__ decides what is saved, an empty dict included, and + # __setstate__ gets it back, as with pickle. Only the default state is + # skipped when empty, see test_container_attrs_saved_only_when_present. + obj = ListWithCustomState([1]) + state = list_get_state(obj, make_save_context()) + assert state["attrs"]["__loader__"] == "DictNode" + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert loaded == obj + assert loaded.restored == {} + + +class SlottedDict(dict): + __slots__ = ("name",) + + +class SlottedList(list): + __slots__ = ("name",) + + +class SlottedSet(set): + __slots__ = ("name",) + + +class SlottedObject: + __slots__ = ("name",) + + +@pytest.mark.skipif( + sys.version_info < (3, 11), + reason="the default __getstate__, which reports the slots, exists since 3.11", +) +@pytest.mark.parametrize( + "slotted_type", [SlottedDict, SlottedList, SlottedSet, SlottedObject] +) +def test_slotted_round_trip(slotted_type): + # The default __getstate__ of an object whose class defines __slots__ + # returns a (dict_state, slots_state) tuple, which is restored the way + # pickle does it, by setting the slots one by one. + obj = slotted_type() + obj.name = "foo" dumped = dumps(obj) loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) - assert type(loaded) is container_type + assert type(loaded) is slotted_type + assert loaded.name == "foo" + + +def test_bunch_round_trip(): + # sklearn's Bunch is a dict subclass which exposes its keys as attributes; + # it is built for pickle's __new__ path and ignores its saved __dict__ in + # its __setstate__, and behaves the same here. + obj = Bunch(a=1, b=[2, 3]) + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded) is Bunch assert loaded == obj + assert loaded.a == 1 + assert loaded.b == [2, 3] + loaded.c = 4 + assert loaded["c"] == 4 -def test_circular_reference_through_list_subclass_raises(): +def test_circular_reference_through_list_subclass(): obj = ListSubclass([1]) obj.append(obj) - msg = "Objects of type ListSubclass which contain a reference to themselves" - with pytest.raises(UnsupportedTypeException, match=msg): - dumps(obj) + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert loaded[0] == 1 + assert loaded[1] is loaded + + +def test_circular_reference_through_dict_subclass(): + obj = DictSubclass(a=1) + obj["self"] = obj + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert loaded["a"] == 1 + assert loaded["self"] is loaded + + +def test_circular_reference_through_defaultdict(): + obj: defaultdict[str, object] = defaultdict(list) + obj["self"] = obj + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded) is defaultdict + assert loaded.default_factory is list + assert loaded["self"] is loaded + + +class DefaultDictWithAttrs(defaultdict): + """defaultdict subclass with a fixed factory and an attribute.""" + + owner: object + + def __init__(self, name): + super().__init__(list) + self.name = name + + +def test_defaultdict_subclass_keeps_type_and_attributes(): + # A defaultdict subclass used to be loaded as a plain defaultdict. Like + # the other container subclasses it is created with __new__, which its + # constructor, unlike with pickle, does not have to accept the factory for. + obj = DefaultDictWithAttrs("foo") + obj["a"].append(1) + obj.owner = obj + dumped = dumps(obj) + assert get_untrusted_types(data=dumped) == [ + f"{get_module(DefaultDictWithAttrs)}.DefaultDictWithAttrs" + ] + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded) is DefaultDictWithAttrs + assert loaded == obj + assert loaded.default_factory is list + assert loaded.name == "foo" + assert loaded.owner is loaded + + +class SetCountingUpdates(set): + """Set subclass whose overridden ``update`` needs its constructor to run.""" + + def __init__(self, items=()): + self.updates = 0 + super().__init__(items) + + def update(self, *args): + self.updates += 1 + super().update(*args) + + +def test_set_subclass_overriding_update_round_trip(): + # The items are added with the built-in set.update, as the constructor + # would, so the override does not run on an instance without attributes. + obj = SetCountingUpdates() + obj.update({1, 2}) + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded) is SetCountingUpdates + assert loaded == obj + assert loaded.updates == 1 + + +def test_defaultdict_non_string_keys(): + # the instance used to be built from the items as keyword arguments + obj = defaultdict(list, {1: [2], 3: [4]}) + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded) is defaultdict + assert loaded.default_factory is list + assert loaded == obj + + +@pytest.mark.parametrize("plain_subclass, with_attrs, items", CONTAINER_SUBCLASS_CASES) +def test_circular_reference_through_container_attribute( + plain_subclass, with_attrs, items +): + obj = with_attrs("foo", items) + obj.owner = obj + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert loaded.owner is loaded + assert loaded.name == "foo" class SetHolder: @@ -1693,11 +1927,13 @@ class SetHolder: members: set -def test_circular_reference_through_set(): +@pytest.mark.parametrize("container_type", [set, SetSubclass]) +def test_circular_reference_through_set(container_type): holder = SetHolder() - holder.members = {holder} + holder.members = container_type({holder}) dumped = dumps(holder) loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded.members) is container_type (member,) = loaded.members assert member is loaded diff --git a/skops/io/tests/test_persist_old.py b/skops/io/tests/test_persist_old.py index 8eaf80e6..6bb7ee67 100644 --- a/skops/io/tests/test_persist_old.py +++ b/skops/io/tests/test_persist_old.py @@ -14,10 +14,22 @@ from skops.io import dumps, get_untrusted_types, loads from skops.io._audit import get_tree -from skops.io._general import OperatorFuncNode, operator_func_get_state +from skops.io._general import ( + DictNode, + ListNode, + OperatorFuncNode, + SetNode, + dict_get_state, + list_get_state, + operator_func_get_state, + set_get_state, +) from skops.io._utils import SaveContext, get_module, get_state, read_schema from skops.io.exceptions import UntrustedTypesFoundException +from skops.io.old._general_v2 import DictNode as DictNodeV2 +from skops.io.old._general_v2 import ListNode as ListNodeV2 from skops.io.old._general_v2 import OperatorFuncNode as OperatorFuncNodeV2 +from skops.io.old._general_v2 import SetNode as SetNodeV2 from skops.io.tests._utils import ( assert_method_outputs_equal, assert_params_equal, @@ -329,10 +341,11 @@ def test_random_generator_v1_wrong_child_type_is_rejected(save_context): ], ids=["attrgetter", "itemgetter", "methodcaller"], ) -def test_operator_func_v2(save_context, func, arg): +@pytest.mark.parametrize("protocol", [0, 1, 2]) +def test_operator_func_v2(save_context, func, arg, protocol): # Up to protocol 2 an OperatorFuncNode state had no "kwargs" entry. Such - # files are read by the protocol-2 node in skops.io.old, and load and - # behave as before. + # files, whichever of these protocols they were written with, are read by + # the protocol-2 node in skops.io.old, and load and behave as before. # operator_func_get_state as it was for protocol 2 def old_operator_func_get_state(obj, save_context): @@ -348,7 +361,7 @@ def old_operator_func_get_state(obj, save_context): data=dumps(func), keys=None, old_state=old_operator_func_get_state(func, save_context), - protocol=2, + protocol=protocol, ) with ZipFile(io.BytesIO(downgraded)) as zip_file: schema, load_context = read_schema(zip_file) @@ -369,3 +382,123 @@ def test_operator_func_current_requires_kwargs(save_context): del state["kwargs"] with pytest.raises(KeyError, match="kwargs"): OperatorFuncNode(state, make_load_context(), trusted=None) + + +class TaggedDict(dict): + """Dict subclass whose constructor sets an attribute.""" + + def __init__(self, items=()): + super().__init__(items) + self.tagged = True + + +class TaggedList(list): + """List subclass whose constructor sets an attribute.""" + + def __init__(self, items): + super().__init__(items) + self.tagged = True + + +class TaggedSet(set): + """Set subclass whose constructor sets an attribute.""" + + def __init__(self, items): + super().__init__(items) + self.tagged = True + + +CONTAINER_SUBCLASS_CASES = [ + pytest.param( + TaggedDict, {"a": 1, "b": 2}, dict_get_state, DictNode, DictNodeV2, id="dict" + ), + pytest.param( + TaggedList, [1, 2, 3], list_get_state, ListNode, ListNodeV2, id="list" + ), + pytest.param(TaggedSet, {1, 2, 3}, set_get_state, SetNode, SetNodeV2, id="set"), +] + + +@pytest.mark.parametrize( + "container_type, items, get_state_func, node_cls, old_node_cls", + CONTAINER_SUBCLASS_CASES, +) +@pytest.mark.parametrize("protocol", [0, 1, 2]) +def test_container_subclass_v2( + save_context, + container_type, + items, + get_state_func, + node_cls, + old_node_cls, + protocol, +): + # Up to protocol 2 the state of a dict, list or set had no "attrs" entry, + # and an instance of a subclass was built through its constructor. Such + # files, whichever of these protocols they were written with, are read by + # the protocol-2 nodes in skops.io.old and load as before, here with the + # attribute the constructor sets. + obj = container_type(items) + # the state as it was for protocol 2 + old_state = get_state_func(obj, save_context) + del old_state["attrs"] + downgraded = downgrade_state( + data=dumps(obj), keys=None, old_state=old_state, protocol=protocol + ) + with ZipFile(io.BytesIO(downgraded)) as zip_file: + schema, load_context = read_schema(zip_file) + node = get_tree(schema, load_context, trusted=None) + assert isinstance(node, old_node_cls) + + type_name = f"{get_module(container_type)}.{container_type.__name__}" + assert get_untrusted_types(data=downgraded) == [type_name] + loaded = loads(downgraded, trusted=[type_name]) + assert type(loaded) is container_type + assert loaded == obj + assert loaded.tagged is True + + +@pytest.mark.parametrize( + "obj, get_state_func, old_node_cls", + [ + pytest.param({"a": 1, "b": 2}, dict_get_state, DictNodeV2, id="dict"), + pytest.param([1, 2, 3], list_get_state, ListNodeV2, id="list"), + pytest.param({1, 2, 3}, set_get_state, SetNodeV2, id="set"), + ], +) +@pytest.mark.parametrize("protocol", [0, 1, 2]) +def test_plain_container_v2(save_context, obj, get_state_func, old_node_cls, protocol): + # A plain dict, list or set has no "attrs" entry in any protocol. Files up + # to protocol 2 are read through the old nodes, which build it as before. + old_state = get_state_func(obj, save_context) + assert "attrs" not in old_state + downgraded = downgrade_state( + data=dumps(obj), keys=None, old_state=old_state, protocol=protocol + ) + with ZipFile(io.BytesIO(downgraded)) as zip_file: + schema, load_context = read_schema(zip_file) + node = get_tree(schema, load_context, trusted=None) + assert isinstance(node, old_node_cls) + loaded = loads(downgraded) + assert type(loaded) is type(obj) + assert loaded == obj + + +@pytest.mark.parametrize( + "container_type, items, get_state_func, node_cls, old_node_cls", + CONTAINER_SUBCLASS_CASES, +) +def test_container_subclass_current_does_not_call_constructor( + save_context, container_type, items, get_state_func, node_cls, old_node_cls +): + # The current nodes create the instance with __new__ and restore its + # attributes from the "attrs" entry added in protocol 3. Without the entry + # the constructor is not called either, which is why a protocol-2 state + # goes through the old node instead, see test_container_subclass_v2. + obj = container_type(items) + state = get_state_func(obj, save_context) + del state["attrs"] + loaded = node_cls(state, make_load_context(), trusted=None).construct() + assert type(loaded) is container_type + assert loaded == obj + assert not hasattr(loaded, "tagged")