Skip to content
Merged
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
11 changes: 11 additions & 0 deletions docs/changes.rst
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,17 @@ 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`_.
- 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
Expand Down
46 changes: 42 additions & 4 deletions skops/io/_audit.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,6 +144,8 @@ def __init__(
self._is_safe = None
# the constructed object can be anything, hence ``Any``
self._constructed: Any = UNINITIALIZED
# set while ``_construct`` runs, see ``construct``
self._constructing = False
saved_id = state.get("__id__")
if saved_id and memoize:
# hold reference to obj in case same instance encountered again in
Expand All @@ -168,10 +170,33 @@ def construct(self) -> Any:
"""Construct the object.

We only construct the object once, and then cache the result.

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
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.
"""
if self._constructed is not UNINITIALIZED:
return self._constructed
self._constructed = self._construct()
if self._constructing:
raise ValueError(
f"Cannot construct an object of type {self.format()} which contains"
" a reference to itself. This is only supported for objects,"
" dicts, lists, and sets."
)

self._constructing = True
try:
self._constructed = self._construct()
except BaseException:
# ``_construct`` may have stored a partially constructed instance
self._constructed = UNINITIALIZED
raise
finally:
self._constructing = False
return self._constructed

def _construct(self) -> Any:
Expand Down Expand Up @@ -306,9 +331,6 @@ def __init__(
self.children = {}

def _construct(self):
# TODO: FIXME This causes a recursion error when loading a cached
# object if we call the cached object's `construct``. Some refactoring
# is needed to fix this.
return self.cached.construct()


Expand Down Expand Up @@ -382,6 +404,22 @@ def get_tree(
# node's ``construct`` method caches the instance.
loaded_tree = load_context.get_object(saved_id)
_check_node_type(type(loaded_tree), allowed_types)
# Within one dump an __id__ belongs to a single object, so a second
# state with the same __id__ is either that object saved again or a
# reference back to it, and both name the same type. Two different
# types sharing an __id__ means the file was not produced by dumping an
# object; refuse it rather than loading one node in place of the other.
class_name, module_name = state.get("__class__"), state.get("__module__")
if (class_name, module_name) != (
loaded_tree.class_name,
loaded_tree.module_name,
):
raise ValueError(
f"The object id {saved_id!r} is used for an object of type"
f" {loaded_tree.module_name}.{loaded_tree.class_name} and for an"
f" object of type {module_name}.{class_name}. This is probably due"
" to a corrupted or a malicious file."
)
return loaded_tree

loader: str = state["__loader__"]
Expand Down
40 changes: 37 additions & 3 deletions skops/io/_general.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,9 @@ def __init__(

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()
Expand Down Expand Up @@ -156,7 +159,19 @@ def __init__(

def _construct(self):
content_type = gettype(self.module_name, self.class_name)
return content_type([item.construct() for item in self.content])
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


def set_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]:
Expand Down Expand Up @@ -186,7 +201,19 @@ def __init__(

def _construct(self):
content_type = gettype(self.module_name, self.class_name)
return content_type([item.construct() for item in self.content])
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


def tuple_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]:
Expand Down Expand Up @@ -548,6 +575,9 @@ def _construct(self):
# issue of required init arguments. Note that the instance created here
# might not be valid until all its attributes have been set below.
instance = cls.__new__(cls)
# 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:
Expand All @@ -571,9 +601,13 @@ def method_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]:
# and prepares both to be persisted.
owner = obj.__self__
func_name = obj.__func__.__name__
# MethodNode checks these against the type of the persisted owner, so the
# module has to be the one of the owner's class, not the one where the
# method is defined, which differs when the method is inherited from a
# class in another module.
res = {
"__class__": owner.__class__.__name__,
"__module__": get_module(obj),
"__module__": get_module(type(owner)),
"__loader__": "MethodNode",
"content": {
"func": func_name,
Expand Down
3 changes: 1 addition & 2 deletions skops/io/_sklearn.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@

from typing import Any

from sklearn.cluster import Birch
from sklearn.tree._tree import Tree

from ._audit import Node, get_tree
Expand Down Expand Up @@ -104,7 +103,7 @@
LossFunction = None


UNSUPPORTED_TYPES = {Birch}
UNSUPPORTED_TYPES: set[type] = set()


def reduce_get_state(obj: Any, save_context: SaveContext) -> dict[str, Any]:
Expand Down
57 changes: 46 additions & 11 deletions skops/io/_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
from zipfile import ZipFile

from ._protocol import PROTOCOL
from .exceptions import UnsupportedTypeException

# Types the user trusts to be loaded, given either by their fully qualified
# name, e.g. "sklearn.linear_model._logistic.LogisticRegression" as returned by
Expand Down Expand Up @@ -128,6 +129,11 @@ class SaveContext:
zip_file: ZipFile
protocol: int = PROTOCOL
memo: dict[int, Any] = field(default_factory=dict)
# The ids of the objects whose state is being computed right now, i.e. the
# path from the root object to the current one, mapped to whether a
# reference back to that object was found inside its own state. Used by
# ``get_state`` to detect circular references.
in_progress: dict[int, bool] = field(default_factory=dict)

def memoize(self, obj: Any) -> int:
# Currently, the only purpose for saving the object id is to make sure
Expand Down Expand Up @@ -221,23 +227,52 @@ 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.

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")


def get_state(value, save_context: SaveContext) -> dict[str, Any]:
# This is a helper function to try to get the state of an object. If it
# fails with `get_state`, we try with json.dumps, if that fails, we raise
# the original error alongside the json error.

# TODO: This should help with fixing recursive references.
# if id(value) in save_context.memo:
# return {
# "__module__": None,
# "__class__": None,
# "__id__": id(value),
# "__loader__": "CachedNode",
# }

__id__ = save_context.memoize(obj=value)

res = _get_state(value, save_context)
if __id__ in save_context.in_progress:
# We are already computing the state of ``value``, so it contains a
# reference to itself, directly or through its children. Instead of
# recursing forever, save a reference to it. When loading, ``get_tree``
# resolves the ``__id__`` to the node of the occurrence which holds the
# actual state, since that node is created before its children.
save_context.in_progress[__id__] = True
# same type description as the state of the object itself, which
# ``get_tree`` checks when it resolves the reference
return {
"__class__": value.__class__.__name__,
"__module__": get_module(type(value)),
"__loader__": "CachedNode",
"__id__": __id__,
}
Comment thread
adrinjalali marked this conversation as resolved.

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
):
raise UnsupportedTypeException(
f"Objects of type {type(value).__name__} which contain a"
" reference to themselves are not supported yet."
)
finally:
del save_context.in_progress[__id__]

res["__id__"] = __id__
return res
Expand Down
14 changes: 12 additions & 2 deletions skops/io/_visualize.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,7 @@ def walk_tree(
node_name: str = "root",
level: int = 0,
is_last: bool = False,
_ancestors: frozenset[int] = frozenset(),
) -> Iterator[NodeInfo]:
"""Visit all nodes of the tree and yield their important attributes.

Expand Down Expand Up @@ -233,6 +234,11 @@ def walk_tree(
is_last: bool (default=False)
Whether this is the last node among its sibling nodes.

_ancestors: frozenset of int (default=frozenset())
The ids of the nodes on the path from the root to this node. A node
which is its own ancestor is a circular reference; it is shown once
more, but its children are not visited again.

Yields
------
:class:`~NodeInfo`:
Expand All @@ -258,6 +264,7 @@ def walk_tree(
node_name=key,
level=level,
is_last=i == num_nodes,
_ancestors=_ancestors,
)
return

Expand All @@ -269,6 +276,7 @@ def walk_tree(
node_name=node_name,
level=level,
is_last=i == num_nodes,
_ancestors=_ancestors,
)
return

Expand All @@ -284,10 +292,11 @@ def walk_tree(
# Note: calling node.is_safe() on all nodes is potentially wasteful because
# it is already a recursive call, i.e. child nodes will be checked many
# times. A solution to this would be to add caching to its call.
is_circular = id(node) in _ancestors
yield NodeInfo(
level=level,
key=node_name,
val=node.format(),
val=node.format() + (" (circular reference)" if is_circular else ""),
is_self_safe=node.is_self_safe(),
is_safe=node.is_safe(),
is_last=is_last,
Expand All @@ -297,13 +306,14 @@ def walk_tree(
# TODO: For better security, we should check the schema if we return early,
# otherwise something nefarious could be hidden inside (however, if there
# is, the node should be marked as unsafe)
if isinstance(node, SKIPPED_TYPES):
if isinstance(node, SKIPPED_TYPES) or is_circular:
return

yield from walk_tree(
node.children,
node_name=node_name,
level=level + 1,
_ancestors=_ancestors | {id(node)},
)


Expand Down
Loading
Loading