From 8a7f867f13919d3ae2aa58179bffe2988985157f Mon Sep 17 00:00:00 2001 From: adrinjalali Date: Sat, 26 Sep 2026 17:36:33 +0100 Subject: [PATCH 1/4] FEAT support circular pointers --- docs/changes.rst | 9 ++ skops/io/_audit.py | 30 ++++++- skops/io/_general.py | 32 ++++++- skops/io/_sklearn.py | 3 +- skops/io/_utils.py | 55 +++++++++--- skops/io/_visualize.py | 14 ++- skops/io/tests/_utils.py | 30 +++++++ skops/io/tests/test_audit.py | 42 +++++++++ skops/io/tests/test_persist.py | 144 +++++++++++++++++++++++++++++-- skops/io/tests/test_visualize.py | 16 ++++ 10 files changed, 347 insertions(+), 28 deletions(-) diff --git a/docs/changes.rst b/docs/changes.rst index 5ff4a358..09f681fc 100644 --- a/docs/changes.rst +++ b/docs/changes.rst @@ -11,6 +11,15 @@ skops Changelog v0.16 ----- +- 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. :issue:`184` by `Adrin Jalali`_. - Fix loading of time-zone-aware ``datetime.datetime`` and ``datetime.time`` objects, and of ``zoneinfo.ZoneInfo`` and ``datetime.timezone`` instances. Their ``__reduce__`` output was not recognized as a constructor call, so they diff --git a/skops/io/_audit.py b/skops/io/_audit.py index 92c2d317..b6bdf0fd 100644 --- a/skops/io/_audit.py +++ b/skops/io/_audit.py @@ -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 @@ -166,10 +168,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: @@ -304,9 +329,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() diff --git a/skops/io/_general.py b/skops/io/_general.py index ef557ec6..6d10fa19 100644 --- a/skops/io/_general.py +++ b/skops/io/_general.py @@ -79,6 +79,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() @@ -149,7 +152,15 @@ 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: + 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]: @@ -179,7 +190,15 @@ 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: + 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]: @@ -520,6 +539,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: @@ -543,9 +565,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, diff --git a/skops/io/_sklearn.py b/skops/io/_sklearn.py index bdc8acec..0d6ee7ab 100644 --- a/skops/io/_sklearn.py +++ b/skops/io/_sklearn.py @@ -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 @@ -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]: diff --git a/skops/io/_utils.py b/skops/io/_utils.py index f5ef986a..195a8293 100644 --- a/skops/io/_utils.py +++ b/skops/io/_utils.py @@ -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 @@ -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 @@ -221,23 +227,50 @@ 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 + return { + "__class__": type(value).__name__, + "__module__": get_module(type(value)), + "__loader__": "CachedNode", + "__id__": __id__, + } + + 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 diff --git a/skops/io/_visualize.py b/skops/io/_visualize.py index 715b29a0..cdcaba19 100644 --- a/skops/io/_visualize.py +++ b/skops/io/_visualize.py @@ -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. @@ -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`: @@ -258,6 +264,7 @@ def walk_tree( node_name=key, level=level, is_last=i == num_nodes, + _ancestors=_ancestors, ) return @@ -269,6 +276,7 @@ def walk_tree( node_name=node_name, level=level, is_last=i == num_nodes, + _ancestors=_ancestors, ) return @@ -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, @@ -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)}, ) diff --git a/skops/io/tests/_utils.py b/skops/io/tests/_utils.py index da3ab675..534cc4fb 100644 --- a/skops/io/tests/_utils.py +++ b/skops/io/tests/_utils.py @@ -4,6 +4,7 @@ import json import sys import warnings +from functools import wraps from zipfile import ZipFile import numpy as np @@ -47,6 +48,33 @@ def _is_steps_like(obj): return True +# The pairs of values being compared further up the stack, see +# ``_skip_circular_references``. +_COMPARING: set[tuple[int, int]] = set() + + +def _skip_circular_references(func): + """Return early for a pair of values which is already being compared. + + Objects can refer back to themselves, directly or through their + attributes, e.g. the tree of a fitted Birch. Comparing such a pair again + would recurse forever; it is being compared further up the stack already. + """ + + @wraps(func) + def wrapper(val1, val2, path=""): + key = (id(val1), id(val2)) + if key in _COMPARING: + return + _COMPARING.add(key) + try: + return func(val1, val2, path=path) + finally: + _COMPARING.discard(key) + + return wrapper + + def _assert_generic_objects_equal(val1, val2, path=""): def _is_builtin(val): # Check if value is a builtin type @@ -76,6 +104,7 @@ def _assert_tuples_equal(val1, val2, path=""): _assert_vals_equal(subval1, subval2, path=f"{path}[]") +@_skip_circular_references def _assert_vals_equal(val1, val2, path=""): if isinstance(val1, type): # e.g. could be np.int64 assert val1 is val2, f"Path: {path}" @@ -155,6 +184,7 @@ def _clean_params(params): return params +@_skip_circular_references def assert_params_equal(params1, params2, path=""): # helper function to compare estimator dictionaries of parameters if params1 is None and params2 is None: diff --git a/skops/io/tests/test_audit.py b/skops/io/tests/test_audit.py index 609aebee..e4f475c9 100644 --- a/skops/io/tests/test_audit.py +++ b/skops/io/tests/test_audit.py @@ -21,6 +21,7 @@ from skops.io._general import ( DictNode, JsonNode, + ListNode, MethodNode, ObjectNode, OperatorFuncNode, @@ -343,3 +344,44 @@ def test_cached_node_resolves_to_memoized_node(): # get_tree short-circuits on an already memoized __id__ and hands back the # original node instead of building a CachedNode. assert get_tree(cached_state, load_context, trusted=None) is node + + +def test_circular_reference_resolves_to_same_node(): + # A reference back to an object whose state is being saved is stored as a + # CachedNode with the __id__ of that object, which get_tree resolves to the + # node holding the object's actual state. + obj: list[object] = [1] + obj.append(obj) + state = get_state(obj, make_save_context()) + assert state["content"][1]["__loader__"] == "CachedNode" + assert state["content"][1]["__id__"] == state["__id__"] + + node = get_tree(state, make_load_context(), trusted=None) + assert isinstance(node, ListNode) + assert node.content[1] is node + assert node.get_unsafe_set() == set() + loaded = node.construct() + assert loaded[1] is loaded + + +def test_construct_refuses_unresolvable_circular_reference(): + # Only nodes which can hand out a partially constructed instance resolve a + # reference back to themselves. For any other node, e.g. a tuple, a file + # claiming such a reference is refused with a clear error instead of + # recursing until the interpreter gives up. dumps never produces such a + # file, so this only happens for a corrupted or malicious one. + state = get_state((1, 2), make_save_context()) + state["content"] = [ + state["content"][0], + { + "__class__": "tuple", + "__module__": "builtins", + "__loader__": "CachedNode", + "__id__": state["__id__"], + }, + ] + node = get_tree(state, make_load_context(), trusted=None) + # the audit terminates, and a tuple of ints is trusted + assert node.get_unsafe_set() == set() + with pytest.raises(ValueError, match="contains a reference to itself"): + node.construct() diff --git a/skops/io/tests/test_persist.py b/skops/io/tests/test_persist.py index 387a1b1f..c1b6afd0 100644 --- a/skops/io/tests/test_persist.py +++ b/skops/io/tests/test_persist.py @@ -17,7 +17,7 @@ import numpy as np import pytest import sklearn -from scipy import sparse, special +from scipy import sparse, special, stats from sklearn.base import BaseEstimator, is_regressor from sklearn.compose import ColumnTransformer from sklearn.datasets import load_sample_images, make_classification, make_regression @@ -1096,14 +1096,16 @@ def test_works_when_given_multiple_bound_methods_attached_to_single_instance(sel loaded_1 = loaded_transformer.inverse_func.__self__ assert loaded_0 is loaded_1 - @pytest.mark.xfail(reason="Failing due to circular self reference", strict=True) - def test_scipy_stats(self, tmp_path): - from scipy import stats - + def test_scipy_stats(self): + # Regression test for gh-184: scipy's discrete distributions keep a + # numpy.vectorize of one of their own bound methods as an attribute, + # which is a reference back to the distribution. estimator = FunctionTransformer(func=stats.zipf) dumped = dumps(estimator) untrusted_types = get_untrusted_types(data=dumped) - loads(dumped, trusted=untrusted_types) + loaded = loads(dumped, trusted=untrusted_types) + assert type(loaded.func) is type(stats.zipf) + assert loaded.func.pmf(2, 3) == pytest.approx(stats.zipf.pmf(2, 3)) class CustomEstimator(BaseEstimator): @@ -1561,3 +1563,133 @@ def fail_gettype(*args, **kwargs): with pytest.raises(UntrustedTypesFoundException, match="malicious_mod.Payload"): loads(dumped) + + +class CircularReferenceEstimator(BaseEstimator): + """Estimator whose fitted attributes refer back to it, see gh-184.""" + + def fit(self, X, y=None): + self.list_: list[object] = [123, self, 456] + self.dict_ = {"a": self.list_} + self.list_.append(self.dict_) + self.method_ = self.meth + return self + + def meth(self): + return "called" + + +def test_circular_references_are_persisted(): + # Regression test for gh-184: objects containing references to themselves, + # directly, through containers, or through bound methods, used to raise a + # RecursionError when dumped. + estimator = CircularReferenceEstimator().fit(None) + dumped = dumps(estimator) + untrusted_types = get_untrusted_types(data=dumped) + loaded = loads(dumped, trusted=untrusted_types) + + assert loaded.list_[0] == 123 + assert loaded.list_[1] is loaded + assert loaded.list_[2] == 456 + assert loaded.list_[3] is loaded.dict_ + assert loaded.dict_["a"] is loaded.list_ + assert loaded.method_.__self__ is loaded + assert loaded.method_() == "called" + + +@pytest.mark.parametrize("dict_type", [dict, OrderedDict]) +def test_circular_reference_in_dict(dict_type): + obj = dict_type(a=1) + obj["self"] = obj + loaded = loads(dumps(obj)) + assert type(loaded) is dict_type + assert loaded["a"] == 1 + assert loaded["self"] is loaded + + +def test_circular_reference_in_list(): + obj: list[object] = [1] + obj.append(obj) + loaded = loads(dumps(obj)) + assert loaded[0] == 1 + assert loaded[1] is loaded + + +class SetHolder: + """Object which is a member of one of its own attributes.""" + + members: set + + +def test_circular_reference_through_set(): + holder = SetHolder() + holder.members = {holder} + dumped = dumps(holder) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + (member,) = loaded.members + assert member is loaded + + +def _tuple_referencing_itself(): + obj = ([],) + obj[0].append(obj) + return obj + + +def _object_array_referencing_itself(): + obj = np.empty(1, dtype=object) + obj[0] = obj + return obj + + +@pytest.mark.parametrize( + "make_obj, type_name", + [ + (_tuple_referencing_itself, "tuple"), + (_object_array_referencing_itself, "ndarray"), + ], +) +def test_circular_reference_through_unsupported_type_raises(make_obj, type_name): + # A tuple or an array cannot be handed out before its items are + # constructed, so a reference to it from one of its items cannot be + # reconstructed. Refuse it when saving instead of producing a file that + # cannot be loaded. + msg = f"Objects of type {type_name} which contain a reference to themselves" + with pytest.raises(UnsupportedTypeException, match=msg): + dumps(make_obj()) + + +def test_shared_references_are_saved_in_full(): + # Only a reference back to an object whose state is being saved is stored + # as a reference. An object which appears several times without referring + # to itself is saved in full each time, as before, so that the file format + # of such objects does not change. Loading still resolves them to the same + # instance. + shared = [1, 2] + dumped = dumps({"a": shared, "b": shared}) + with ZipFile(io.BytesIO(dumped), "r") as zip_file: + schema = zip_file.read("schema.json").decode() + assert "CachedNode" not in schema + + loaded = loads(dumped) + assert loaded["a"] == [1, 2] + assert loaded["a"] is loaded["b"] + + +class EstimatorInheritingMethods(BaseEstimator): + """Its methods are defined in ``sklearn.base``, not in this module.""" + + +def test_method_defined_in_other_module(): + # The module saved for a bound method is the module of the owner's class, + # which MethodNode checks the file against, not the module where the method + # is defined. Files of inherited methods used to be rejected as corrupted. + method = EstimatorInheritingMethods().get_params + state = get_state(method, make_save_context()) + assert state["__module__"] == __name__ + + dumped = dumps(method) + untrusted_types = get_untrusted_types(data=dumped) + assert f"{__name__}.EstimatorInheritingMethods.get_params" in untrusted_types + loaded = loads(dumped, trusted=untrusted_types) + assert loaded() == method() diff --git a/skops/io/tests/test_visualize.py b/skops/io/tests/test_visualize.py index ef470b48..e2b1640f 100644 --- a/skops/io/tests/test_visualize.py +++ b/skops/io/tests/test_visualize.py @@ -338,3 +338,19 @@ def test_decision_tree(self, cls, capsys): assert expected_tree_block in stdout assert " │ └── constructor: sklearn.tree._tree.Tree [UNSAFE]" in stdout assert '_sklearn_version: json-type("{}")'.format(sklearn.__version__) in stdout + + +def test_visualize_circular_reference(capsys): + # A node which contains itself is shown once more, marked as such, and its + # children are not visited again. + obj: list[object] = [1] + obj.append(obj) + sio.visualize(sio.dumps(obj)) + + expected = [ + "root: builtins.list", + "├── content: json-type(1)", + "└── content: builtins.list (circular reference)", + ] + stdout, _ = capsys.readouterr() + assert stdout.strip() == "\n".join(expected) From 84f5d592447ea5598f3da6bef79bfc15e629046a Mon Sep 17 00:00:00 2001 From: adrinjalali Date: Sun, 27 Sep 2026 08:39:41 +0100 Subject: [PATCH 2/4] harden --- docs/changes.rst | 4 ++- skops/io/_audit.py | 16 +++++++++++- skops/io/_utils.py | 4 ++- skops/io/tests/_utils.py | 28 ++++++++++++++++++-- skops/io/tests/test_audit.py | 17 +++++++++++++ skops/io/tests/test_persist.py | 27 ++++++++++++++++++++ skops/io/tests/test_persist_old.py | 41 ++++++++++++++++++++++++++++++ 7 files changed, 132 insertions(+), 5 deletions(-) diff --git a/docs/changes.rst b/docs/changes.rst index 09f681fc..d6485446 100644 --- a/docs/changes.rst +++ b/docs/changes.rst @@ -19,7 +19,9 @@ v0.16 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. :issue:`184` by `Adrin Jalali`_. + 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`_. - Fix loading of time-zone-aware ``datetime.datetime`` and ``datetime.time`` objects, and of ``zoneinfo.ZoneInfo`` and ``datetime.timezone`` instances. Their ``__reduce__`` output was not recognized as a constructor call, so they diff --git a/skops/io/_audit.py b/skops/io/_audit.py index b6bdf0fd..b15682c7 100644 --- a/skops/io/_audit.py +++ b/skops/io/_audit.py @@ -375,7 +375,21 @@ def get_tree( # the parent node's ``construct`` method is called, and for this node # it'll be called more than once. But that's not an issue since the # node's ``construct`` method caches the instance. - return load_context.get_object(saved_id) + node = load_context.get_object(saved_id) + # 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) != (node.class_name, node.module_name): + raise ValueError( + f"The object id {saved_id!r} is used for an object of type" + f" {node.module_name}.{node.class_name} and for an object of type" + f" {module_name}.{class_name}. This is probably due to a corrupted" + " or a malicious file." + ) + return node loader: str = state["__loader__"] protocol = load_context.protocol diff --git a/skops/io/_utils.py b/skops/io/_utils.py index 195a8293..699bb473 100644 --- a/skops/io/_utils.py +++ b/skops/io/_utils.py @@ -252,8 +252,10 @@ def get_state(value, save_context: SaveContext) -> dict[str, Any]: # 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__": type(value).__name__, + "__class__": value.__class__.__name__, "__module__": get_module(type(value)), "__loader__": "CachedNode", "__id__": __id__, diff --git a/skops/io/tests/_utils.py b/skops/io/tests/_utils.py index 534cc4fb..28bd72c4 100644 --- a/skops/io/tests/_utils.py +++ b/skops/io/tests/_utils.py @@ -235,6 +235,30 @@ def assert_method_outputs_equal(estimator, loaded, X): assert_allclose_dense_sparse(X_out1, X_out2, err_msg=err_msg, atol=ATOL) +def _unused_id(schema: dict) -> int: + """Return an ``__id__`` which no node of ``schema`` uses. + + The ``__id__`` of a node is the ``id()`` of the object at dump time. Some of + those objects were temporaries which have been freed since, so the ``id()`` + of an object created after dumping can coincide with one of them. Two nodes + sharing an ``__id__`` are loaded as the same object. + """ + used: set[int] = set() + + def collect(state): + if isinstance(state, dict): + if "__id__" in state: + used.add(state["__id__"]) + for value in state.values(): + collect(value) + elif isinstance(state, list): + for value in state: + collect(value) + + collect(schema) + return max(used, default=0) + 1 + + def downgrade_state( *, data: bytes, keys: list[str] | None, old_state: dict, protocol: int ): @@ -302,7 +326,7 @@ def downgrade_state( if keys is None: # replace all fields schema = old_state - schema["__id__"] = id(schema) + schema["__id__"] = _unused_id(schema) else: # replace specific field state = schema @@ -311,7 +335,7 @@ def downgrade_state( state[keys[-1]] = old_state # there has to be an __id__ field for memoization - state[keys[-1]]["__id__"] = id(schema) + state[keys[-1]]["__id__"] = _unused_id(schema) schema["protocol"] = protocol diff --git a/skops/io/tests/test_audit.py b/skops/io/tests/test_audit.py index e4f475c9..b3236843 100644 --- a/skops/io/tests/test_audit.py +++ b/skops/io/tests/test_audit.py @@ -364,6 +364,23 @@ def test_circular_reference_resolves_to_same_node(): assert loaded[1] is loaded +def test_get_tree_rejects_id_shared_by_different_types(): + # An __id__ identifies one object of a dump, so all states carrying it + # describe the same type. A file where a state of another type carries an + # already seen __id__ is refused, instead of loading the first node in its + # place. + state = get_state({"a": [1], "b": (2,)}, make_save_context()) + content = state["content"] + content["b"]["__id__"] = content["a"]["__id__"] + + msg = re.escape( + "is used for an object of type builtins.list and for an object of type" + " builtins.tuple" + ) + with pytest.raises(ValueError, match=msg): + get_tree(state, make_load_context(), trusted=None) + + def test_construct_refuses_unresolvable_circular_reference(): # Only nodes which can hand out a partially constructed instance resolve a # reference back to themselves. For any other node, e.g. a tuple, a file diff --git a/skops/io/tests/test_persist.py b/skops/io/tests/test_persist.py index c1b6afd0..886855ea 100644 --- a/skops/io/tests/test_persist.py +++ b/skops/io/tests/test_persist.py @@ -1615,6 +1615,33 @@ def test_circular_reference_in_list(): assert loaded[1] is loaded +class ListSubclass(list): + pass + + +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]) + dumped = dumps(obj) + loaded = loads(dumped, trusted=get_untrusted_types(data=dumped)) + assert type(loaded) is container_type + assert loaded == obj + + +def test_circular_reference_through_list_subclass_raises(): + 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) + + class SetHolder: """Object which is a member of one of its own attributes.""" diff --git a/skops/io/tests/test_persist_old.py b/skops/io/tests/test_persist_old.py index b32a77c9..341f954a 100644 --- a/skops/io/tests/test_persist_old.py +++ b/skops/io/tests/test_persist_old.py @@ -25,6 +25,47 @@ def dummy_func(X): return X +def test_downgrade_state_assigns_unused_id(): + # The __id__ given to the downgraded node must not be in use by another + # node of the file, otherwise both are loaded as the same object. Using the + # id() of a new object is not enough: the ids in the file belong to objects + # of dump time, some of which have been freed since. + dumped = dumps(FunctionTransformer(func=np.sqrt)) + old_state = { + "__class__": "ufunc", + "__module__": "numpy", + "__loader__": "FunctionNode", + "content": {"module_path": "numpy", "function": "sqrt"}, + } + downgraded = downgrade_state( + data=dumped, + keys=["content", "content", "func"], + old_state=old_state, + protocol=0, + ) + with ZipFile(io.BytesIO(downgraded), "r") as zip_file: + schema = json.loads(zip_file.read("schema.json")) + func_state = schema["content"]["content"]["func"] + + # the ids of all other nodes; shared objects like None legitimately repeat + other_ids: set[int] = set() + + def collect(state): + if state is func_state: + return + if isinstance(state, dict): + if "__id__" in state: + other_ids.add(state["__id__"]) + for value in state.values(): + collect(value) + elif isinstance(state, list): + for value in state: + collect(value) + + collect(schema) + assert func_state["__id__"] not in other_ids + + @pytest.fixture def save_context(): buffer = io.BytesIO() From 4e007afbf7a5bb94a43cd81450dc90b79e505736 Mon Sep 17 00:00:00 2001 From: adrinjalali Date: Sun, 27 Sep 2026 08:40:06 +0100 Subject: [PATCH 3/4] fix changelog --- docs/changes.rst | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/docs/changes.rst b/docs/changes.rst index c36ab768..01128c71 100644 --- a/docs/changes.rst +++ b/docs/changes.rst @@ -20,15 +20,6 @@ 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`_. -- 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 - from the command line since. ``python -m skops`` now also runs the CLI, and - running it without a subcommand prints a usage error instead of a traceback. - :pr:`548` by `Adrin Jalali`_. - -v0.16 ------ - 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 @@ -40,6 +31,15 @@ v0.16 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 + from the command line since. ``python -m skops`` now also runs the CLI, and + running it without a subcommand prints a usage error instead of a traceback. + :pr:`548` by `Adrin Jalali`_. + +v0.16 +----- - Fix loading of time-zone-aware ``datetime.datetime`` and ``datetime.time`` objects, and of ``zoneinfo.ZoneInfo`` and ``datetime.timezone`` instances. Their ``__reduce__`` output was not recognized as a constructor call, so they From 465126b073ea5c9219dba7876c364cefe5b04b99 Mon Sep 17 00:00:00 2001 From: adrinjalali Date: Tue, 29 Sep 2026 08:18:19 +0100 Subject: [PATCH 4/4] comment --- skops/io/_general.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/skops/io/_general.py b/skops/io/_general.py index c635ff2a..7e8c2136 100644 --- a/skops/io/_general.py +++ b/skops/io/_general.py @@ -160,6 +160,10 @@ def __init__( 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 @@ -198,6 +202,10 @@ def __init__( 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