From 66906d777f0ec3df42c35f5740e13985a1086e7b Mon Sep 17 00:00:00 2001 From: adrinjalali Date: Sat, 26 Sep 2026 17:13:47 +0100 Subject: [PATCH 1/3] SEC harden get_tree loaded types --- docs/changes.rst | 9 ++ skops/io/_audit.py | 72 ++++++++++------ skops/io/_general.py | 36 ++++++-- skops/io/_numpy.py | 47 ++++++++-- skops/io/_sklearn.py | 23 ++++- skops/io/old/_numpy_v1.py | 6 +- skops/io/tests/test_audit.py | 134 ++++++++++++++++++++++++++++- skops/io/tests/test_persist_old.py | 24 ++++++ 8 files changed, 305 insertions(+), 46 deletions(-) diff --git a/docs/changes.rst b/docs/changes.rst index 5ff4a358..4eb054c7 100644 --- a/docs/changes.rst +++ b/docs/changes.rst @@ -9,6 +9,15 @@ skops Changelog :depth: 1 :local: +v0.17 +----- +- Loading a skops file now checks that each part of the file holds the kind + of content its loader expects, for instance that the keyword arguments of a + ``functools.partial`` are stored as a dict. A file that does not is refused + with an error while it is read, before anything in it is audited or + constructed, instead of failing with an unrelated error, or being accepted, + during construction. :issue:`222` by `Adrin Jalali`_. + v0.16 ----- - Fix loading of time-zone-aware ``datetime.datetime`` and ``datetime.time`` diff --git a/skops/io/_audit.py b/skops/io/_audit.py index 92c2d317..f64e5cc0 100644 --- a/skops/io/_audit.py +++ b/skops/io/_audit.py @@ -156,6 +156,8 @@ def __init__( # list of appropriate trusted types # 3. store each child in a typed attribute and mirror them in # self.children; do not construct the children objects yet + # 4. pass ``allowed_types`` to ``get_tree`` for each child whose kind + # ``_construct`` relies on, e.g. a dict of keyword arguments self.trusted = self._get_trusted(trusted, []) # ``children`` is the generic view of the node's children, used to audit # and visualize the tree. Subclasses read their children through the @@ -317,6 +319,8 @@ def get_tree( state: dict[str, Any], load_context: LoadContext, trusted: TrustedTypes | None, + *, + allowed_types: tuple[type[Node], ...] | None = None, ) -> Node: """Get the tree of nodes. @@ -341,6 +345,15 @@ def get_tree( objects of types listed in ``trusted`` in the dumped file. Types can be given by their fully qualified name or as the type itself. + allowed_types : tuple of Node subclasses, default=None + If given, the returned node has to be an instance of one of these + classes, and a ``ValueError`` is raised otherwise. A node passes this + for the children whose kind its ``_construct`` relies on, e.g. a + ``PartialNode`` requires its keyword arguments to be a ``DictNode``, + so that a crafted file cannot put an arbitrary node in that place. + The check is done on the node that is returned, so it also covers a + node taken from the memo through its ``__id__``. + Returns ------- loaded_tree : Node @@ -353,30 +366,39 @@ 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) - - loader: str = state["__loader__"] - protocol = load_context.protocol - key = (loader, protocol) - - if key in NODE_TYPE_MAPPING: - node_cls = NODE_TYPE_MAPPING[key] + loaded_tree = load_context.get_object(saved_id) else: - # What probably happened here is that we released a new protocol. If - # there is no specific key for the old protocol, it means it is safe to - # use the current protocol instead, because this node was not changed. - key_new = (loader, PROTOCOL) - try: - node_cls = NODE_TYPE_MAPPING[key_new] - except KeyError: - # If we still cannot find the loader for this key, something went - # wrong. - type_name = f"{state['__module__']}.{state['__class__']}" - raise TypeError( - f" Can't find loader {state['__loader__']} for type {type_name} and " - f"protocol {protocol}. You might need to update skops to load this " - "file." - ) - - loaded_tree = node_cls(state, load_context, trusted=trusted) + loader: str = state["__loader__"] + protocol = load_context.protocol + key = (loader, protocol) + + if key in NODE_TYPE_MAPPING: + node_cls = NODE_TYPE_MAPPING[key] + else: + # What probably happened here is that we released a new protocol. + # If there is no specific key for the old protocol, it means it is + # safe to use the current protocol instead, because this node was + # not changed. + key_new = (loader, PROTOCOL) + try: + node_cls = NODE_TYPE_MAPPING[key_new] + except KeyError: + # If we still cannot find the loader for this key, something + # went wrong. + type_name = f"{state['__module__']}.{state['__class__']}" + raise TypeError( + f" Can't find loader {state['__loader__']} for type {type_name} " + f"and protocol {protocol}. You might need to update skops to " + "load this file." + ) + + loaded_tree = node_cls(state, load_context, trusted=trusted) + + if allowed_types is not None and not isinstance(loaded_tree, allowed_types): + expected = " or ".join(cls.__name__ for cls in allowed_types) + raise ValueError( + f"Expected a node of type {expected}, got " + f"{type(loaded_tree).__name__}. This is probably due to a corrupted " + "or a malicious file." + ) return loaded_tree diff --git a/skops/io/_general.py b/skops/io/_general.py index ef557ec6..ed8ca960 100644 --- a/skops/io/_general.py +++ b/skops/io/_general.py @@ -70,7 +70,9 @@ def __init__( ) -> None: super().__init__(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) + 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() @@ -109,7 +111,12 @@ def __init__( ) -> None: super().__init__(state, load_context, trusted) self.trusted = ["collections.defaultdict"] - self.main = get_tree(state["content"]["main"], load_context, trusted=trusted) + self.main = get_tree( + state["content"]["main"], + load_context, + trusted=trusted, + allowed_types=(DictNode,), + ) self.default_factory = get_tree( state["content"]["default_factory"], load_context, trusted=trusted ) @@ -294,9 +301,19 @@ def __init__( self.trusted = self._get_trusted(trusted, []) content = state["content"] self.func = get_tree(content["func"], load_context, trusted=trusted) - self.args = get_tree(content["args"], load_context, trusted=trusted) - self.kwds = get_tree(content["kwds"], load_context, trusted=trusted) - self.namespace = get_tree(content["namespace"], load_context, trusted=trusted) + self.args = get_tree( + content["args"], load_context, trusted=trusted, allowed_types=(TupleNode,) + ) + self.kwds = get_tree( + content["kwds"], load_context, trusted=trusted, allowed_types=(DictNode,) + ) + # the ``__dict__`` of the partial object, or None if it had none + self.namespace = get_tree( + content["namespace"], + load_context, + trusted=trusted, + allowed_types=(DictNode, JsonNode), + ) self.children = { "func": self.func, "args": self.args, @@ -484,7 +501,10 @@ def __init__( trusted: TrustedTypes | None = None, ) -> None: super().__init__(state, load_context, trusted) - self.content = get_tree(state["content"], load_context, trusted=trusted) + # ``__reduce__`` gives the constructor arguments as a tuple + self.content = get_tree( + state["content"], load_context, trusted=trusted, allowed_types=(TupleNode,) + ) self.children = {"content": self.content} def _construct(self): @@ -713,7 +733,9 @@ def __init__( " due to a corrupted or a malicious file." ) self.trusted = self._get_trusted(trusted, []) - self.attrs = get_tree(state["attrs"], load_context, trusted=trusted) + self.attrs = get_tree( + state["attrs"], load_context, trusted=trusted, allowed_types=(TupleNode,) + ) self.children = {"attrs": self.attrs} def _construct(self): diff --git a/skops/io/_numpy.py b/skops/io/_numpy.py index deaa7192..1b9d6681 100644 --- a/skops/io/_numpy.py +++ b/skops/io/_numpy.py @@ -7,7 +7,7 @@ import numpy as np from ._audit import Node, get_tree -from ._general import function_get_state +from ._general import DictNode, TupleNode, function_get_state from ._protocol import PROTOCOL from ._trusted_types import ( NUMPY_DTYPE_TYPE_NAMES, @@ -86,7 +86,12 @@ def __init__( self.items = [ get_tree(o, load_context, trusted=trusted) for o in state["content"] ] - self.shape = get_tree(state["shape"], load_context, trusted=trusted) + self.shape = get_tree( + state["shape"], + load_context, + trusted=trusted, + allowed_types=(TupleNode,), + ) self.children = {"content": self.items, "shape": self.shape} else: raise ValueError(f"Unknown type {self.type}.") @@ -141,8 +146,20 @@ def __init__( ) -> None: super().__init__(state, load_context, trusted) self.trusted = self._get_trusted(trusted, [np.ma.MaskedArray]) - self.data = get_tree(state["content"]["data"], load_context, trusted=trusted) - self.mask = get_tree(state["content"]["mask"], load_context, trusted=trusted) + # ``mask`` is an array, or the ``numpy.bool_`` scalar ``nomask``; both + # are saved as an ``NdArrayNode`` + self.data = get_tree( + state["content"]["data"], + load_context, + trusted=trusted, + allowed_types=(NdArrayNode,), + ) + self.mask = get_tree( + state["content"]["mask"], + load_context, + trusted=trusted, + allowed_types=(NdArrayNode,), + ) self.children = {"data": self.data, "mask": self.mask} def _construct(self): @@ -171,7 +188,9 @@ def __init__( ) -> None: super().__init__(state, load_context, trusted) # TODO - self.content = get_tree(state["content"], load_context, trusted=trusted) + self.content = get_tree( + state["content"], load_context, trusted=trusted, allowed_types=(DictNode,) + ) self.children = {"content": self.content} self.trusted = self._get_trusted(trusted, [np.random.RandomState]) @@ -202,10 +221,16 @@ def __init__( ) -> None: super().__init__(state, load_context, trusted) self.bit_generator_state = get_tree( - state["content"]["bit_generator"], load_context, trusted=trusted + state["content"]["bit_generator"], + load_context, + trusted=trusted, + allowed_types=(DictNode,), ) self.seed_seq_state = get_tree( - state["content"]["seed_seq"], load_context, trusted=trusted + state["content"]["seed_seq"], + load_context, + trusted=trusted, + allowed_types=(DictNode,), ) self.children = { "bit_generator_state": self.bit_generator_state, @@ -326,7 +351,13 @@ def __init__( trusted: TrustedTypes | None = None, ) -> None: super().__init__(state, load_context, trusted) - self.content = get_tree(state["content"], load_context, trusted=trusted) + # the dtype is stored through an empty array of that dtype + self.content = get_tree( + state["content"], + load_context, + trusted=trusted, + allowed_types=(NdArrayNode,), + ) self.children = {"content": self.content} # TODO: what should we trust? self.trusted = self._get_trusted(trusted, []) diff --git a/skops/io/_sklearn.py b/skops/io/_sklearn.py index bdc8acec..d053a612 100644 --- a/skops/io/_sklearn.py +++ b/skops/io/_sklearn.py @@ -6,7 +6,7 @@ from sklearn.tree._tree import Tree from ._audit import Node, get_tree -from ._general import TypeNode, unsupported_get_state +from ._general import DictNode, TupleNode, TypeNode, unsupported_get_state from ._protocol import PROTOCOL from ._utils import ( LoadContext, @@ -166,8 +166,17 @@ def __init__( super().__init__(state, load_context, trusted) reduce = state["__reduce__"] ctor_module, ctor_class = constructor - self.attrs = get_tree(state["content"], load_context, trusted=trusted) - self.args = get_tree(reduce["args"], load_context, trusted=trusted) + # ``reduce_get_state`` only accepts a dict or a tuple as the state, and + # ``__reduce__`` gives the constructor arguments as a tuple + self.attrs = get_tree( + state["content"], + load_context, + trusted=trusted, + allowed_types=(DictNode, TupleNode), + ) + self.args = get_tree( + reduce["args"], load_context, trusted=trusted, allowed_types=(TupleNode,) + ) self.constructor = TypeNode( {"__class__": ctor_class, "__module__": ctor_module}, load_context, @@ -327,11 +336,17 @@ def __init__( self.trusted = [ get_module(_DictWithDeprecatedKeysNode) + "._DictWithDeprecatedKeys" ] - self.main = get_tree(state["content"]["main"], load_context, trusted=trusted) + self.main = get_tree( + state["content"]["main"], + load_context, + trusted=trusted, + allowed_types=(DictNode,), + ) self.deprecated_key_to_new_key = get_tree( state["content"]["_deprecated_key_to_new_key"], load_context, trusted=trusted, + allowed_types=(DictNode,), ) self.children = { "main": self.main, diff --git a/skops/io/old/_numpy_v1.py b/skops/io/old/_numpy_v1.py index 64b9e30f..861381ac 100644 --- a/skops/io/old/_numpy_v1.py +++ b/skops/io/old/_numpy_v1.py @@ -6,6 +6,7 @@ import numpy as np from skops.io._audit import Node, get_tree +from skops.io._general import DictNode from skops.io._trusted_types import NUMPY_RANDOM_BIT_GENERATOR_TYPE_NAMES from skops.io._utils import LoadContext, TrustedTypes, gettype @@ -21,7 +22,10 @@ def __init__( ) -> None: super().__init__(state, load_context, trusted) self.bit_generator_state = get_tree( - state["content"]["bit_generator"], load_context, trusted=trusted + state["content"]["bit_generator"], + load_context, + trusted=trusted, + allowed_types=(DictNode,), ) self.children = {"bit_generator_state": self.bit_generator_state} self.trusted = self._get_trusted(trusted, [np.random.Generator]) diff --git a/skops/io/tests/test_audit.py b/skops/io/tests/test_audit.py index 609aebee..d7f89869 100644 --- a/skops/io/tests/test_audit.py +++ b/skops/io/tests/test_audit.py @@ -2,14 +2,19 @@ import json import operator import re +from collections import defaultdict from contextlib import suppress +from datetime import timezone +from functools import partial from zipfile import ZipFile +import numpy as np import pytest from sklearn.linear_model import LogisticRegression from sklearn.preprocessing import FunctionTransformer +from sklearn.tree import DecisionTreeClassifier -from skops.io import dumps, get_untrusted_types +from skops.io import dumps, get_untrusted_types, loads from skops.io._audit import ( CachedNode, Node, @@ -21,9 +26,11 @@ from skops.io._general import ( DictNode, JsonNode, + ListNode, MethodNode, ObjectNode, OperatorFuncNode, + TupleNode, dict_get_state, method_get_state, operator_func_get_state, @@ -343,3 +350,128 @@ 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 _replace_child(data: bytes, keys: list[str | int], new_state: dict) -> bytes: + """Return a copy of a skops dump with the node at ``keys`` replaced.""" + src = ZipFile(io.BytesIO(data)) + schema = json.loads(src.read("schema.json")) + parent = schema + for key in keys[:-1]: + parent = parent[key] + parent[keys[-1]] = new_state + buffer = io.BytesIO() + with ZipFile(buffer, "w") as out: + for name in src.namelist(): + if name == "schema.json": + out.writestr(name, json.dumps(schema)) + else: + out.writestr(name, src.read(name)) + return buffer.getvalue() + + +# A node that is trusted by default and that no parent expects as a child. +SLICE_STATE = { + "__class__": "slice", + "__module__": "builtins", + "__loader__": "SliceNode", + "content": {"start": None, "stop": None, "step": None}, +} + +_TREE = DecisionTreeClassifier(random_state=0).fit([[0.0], [1.0]], [0, 1]).tree_ + + +def test_get_tree_allowed_types(): + load_context = make_load_context() + state = get_state((1, 2), make_save_context()) + + node = get_tree(state, load_context, trusted=None, allowed_types=(TupleNode,)) + assert isinstance(node, TupleNode) + assert ( + get_tree(state, load_context, trusted=None, allowed_types=(DictNode, TupleNode)) + is node + ) + + msg = "Expected a node of type DictNode or ListNode, got TupleNode" + with pytest.raises(ValueError, match=msg): + get_tree( + state, + make_load_context(), + trusted=None, + allowed_types=(DictNode, ListNode), + ) + # The node is in the memo of ``load_context`` by now and is handed back + # from there; the check applies to that node too. + with pytest.raises(ValueError, match=msg): + get_tree(state, load_context, trusted=None, allowed_types=(DictNode, ListNode)) + + +@pytest.mark.parametrize( + "obj, keys, expected", + [ + ({"a": 1}, ["key_types"], "ListNode"), + (defaultdict(list), ["content", "main"], "DictNode"), + (partial(np.add, 1), ["content", "args"], "TupleNode"), + (partial(np.add, 1), ["content", "kwds"], "DictNode"), + (partial(np.add, 1), ["content", "namespace"], "DictNode or JsonNode"), + (operator.itemgetter(1), ["attrs"], "TupleNode"), + (timezone.utc, ["content"], "TupleNode"), + (np.array([1, "a"], dtype=object), ["shape"], "TupleNode"), + (np.ma.MaskedArray([1, 2]), ["content", "mask"], "NdArrayNode"), + (np.random.RandomState(0), ["content"], "DictNode"), + (np.random.default_rng(0), ["content", "seed_seq"], "DictNode"), + (np.dtype("float64"), ["content"], "NdArrayNode"), + (_TREE, ["content"], "DictNode or TupleNode"), + (_TREE, ["__reduce__", "args"], "TupleNode"), + ], + ids=[ + "dict.key_types", + "defaultdict.main", + "partial.args", + "partial.kwds", + "partial.namespace", + "itemgetter.attrs", + "timezone.content", + "ndarray.shape", + "maskedarray.mask", + "randomstate.content", + "generator.seed_seq", + "dtype.content", + "tree.attrs", + "tree.args", + ], +) +def test_child_of_unexpected_node_type_is_rejected(obj, keys, expected): + # Each node states which kind of child its ``_construct`` relies on, and a + # file with any other node there is refused while the tree is built, i.e. + # before the audit. So neither loading nor listing the untrusted types + # gets as far as constructing anything from such a file, and the user is + # told that the file is corrupted instead of getting an unrelated error + # from construct, or no error at all. + data = _replace_child(dumps(obj), keys, SLICE_STATE) + msg = f"Expected a node of type {expected}, got SliceNode" + with pytest.raises(ValueError, match=msg): + get_untrusted_types(data=data) + with pytest.raises(ValueError, match=msg): + loads(data) + + +def test_child_aliased_to_earlier_node_is_rejected(): + # ``get_tree`` hands back the node it already built for an ``__id__`` it + # has seen before, whatever else the state says, so a file can point a + # child at any earlier node of the tree. The type is checked on the node + # that is handed back, which catches this too. + data = dumps([LogisticRegression(), {"a": 1}]) + schema = json.loads(ZipFile(io.BytesIO(data)).read("schema.json")) + alias = { + "__class__": "list", + "__module__": "builtins", + "__loader__": "ListNode", + "__id__": schema["content"][0]["__id__"], + } + data = _replace_child(data, ["content", 1, "key_types"], alias) + msg = "Expected a node of type ListNode, got ObjectNode" + with pytest.raises(ValueError, match=msg): + get_untrusted_types(data=data) + with pytest.raises(ValueError, match=msg): + loads(data) diff --git a/skops/io/tests/test_persist_old.py b/skops/io/tests/test_persist_old.py index b32a77c9..818ecd92 100644 --- a/skops/io/tests/test_persist_old.py +++ b/skops/io/tests/test_persist_old.py @@ -243,3 +243,27 @@ def test_random_generator_v1_missing_name_is_rejected(save_context): ) with pytest.raises(ValueError, match="Could not find the bit generator name"): get_untrusted_types(data=broken) + + +def test_random_generator_v1_wrong_child_type_is_rejected(save_context): + # As for the current node (see test_audit.py), the bit generator state has + # to be a DictNode, and a file with any other node there is refused before + # the audit. + rng = np.random.default_rng(42) + slice_state = { + "__class__": "slice", + "__module__": "builtins", + "__loader__": "SliceNode", + "content": {"start": None, "stop": None, "step": None}, + } + broken = downgrade_state( + data=_dump_v1_generator(save_context, rng), + keys=["content", "bit_generator"], + old_state=slice_state, + protocol=1, + ) + msg = "Expected a node of type DictNode, got SliceNode" + with pytest.raises(ValueError, match=msg): + get_untrusted_types(data=broken) + with pytest.raises(ValueError, match=msg): + loads(broken) From 5e65624cd155d9c126d29d2b633b3c9c966e6b97 Mon Sep 17 00:00:00 2001 From: adrinjalali Date: Sat, 26 Sep 2026 17:19:33 +0100 Subject: [PATCH 2/3] changelog --- docs/changes.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/changes.rst b/docs/changes.rst index 4eb054c7..d7138fda 100644 --- a/docs/changes.rst +++ b/docs/changes.rst @@ -16,7 +16,7 @@ v0.17 ``functools.partial`` are stored as a dict. A file that does not is refused with an error while it is read, before anything in it is audited or constructed, instead of failing with an unrelated error, or being accepted, - during construction. :issue:`222` by `Adrin Jalali`_. + during construction. :pr:`547` by `Adrin Jalali`_. v0.16 ----- From 0fb4502d8907934470c1c1f69b68839432bb67f9 Mon Sep 17 00:00:00 2001 From: adrinjalali Date: Sun, 27 Sep 2026 07:49:52 +0100 Subject: [PATCH 3/3] fixes --- skops/io/_audit.py | 81 ++++++++++++++++++++---------------- skops/io/tests/test_audit.py | 22 ++++++++++ 2 files changed, 67 insertions(+), 36 deletions(-) diff --git a/skops/io/_audit.py b/skops/io/_audit.py index f64e5cc0..c866f2b5 100644 --- a/skops/io/_audit.py +++ b/skops/io/_audit.py @@ -315,6 +315,19 @@ def _construct(self): NODE_TYPE_MAPPING[("CachedNode", PROTOCOL)] = CachedNode +def _check_node_type( + node_cls: type[Node], allowed_types: tuple[type[Node], ...] | None +) -> None: + """Raise ``ValueError`` if ``node_cls`` is not one of ``allowed_types``.""" + if allowed_types is None or issubclass(node_cls, allowed_types): + return + expected = " or ".join(cls.__name__ for cls in allowed_types) + raise ValueError( + f"Expected a node of type {expected}, got {node_cls.__name__}. This is " + "probably due to a corrupted or a malicious file." + ) + + def get_tree( state: dict[str, Any], load_context: LoadContext, @@ -351,8 +364,9 @@ def get_tree( for the children whose kind its ``_construct`` relies on, e.g. a ``PartialNode`` requires its keyword arguments to be a ``DictNode``, so that a crafted file cannot put an arbitrary node in that place. - The check is done on the node that is returned, so it also covers a - node taken from the memo through its ``__id__``. + The loader named in ``state`` is checked before its node is built, so + nothing in a rejected subtree is read, and a node handed back from the + memo through its ``__id__`` is checked as well. Returns ------- @@ -367,38 +381,33 @@ def get_tree( # it'll be called more than once. But that's not an issue since the # node's ``construct`` method caches the instance. loaded_tree = load_context.get_object(saved_id) + _check_node_type(type(loaded_tree), allowed_types) + return loaded_tree + + loader: str = state["__loader__"] + protocol = load_context.protocol + key = (loader, protocol) + + if key in NODE_TYPE_MAPPING: + node_cls = NODE_TYPE_MAPPING[key] else: - loader: str = state["__loader__"] - protocol = load_context.protocol - key = (loader, protocol) - - if key in NODE_TYPE_MAPPING: - node_cls = NODE_TYPE_MAPPING[key] - else: - # What probably happened here is that we released a new protocol. - # If there is no specific key for the old protocol, it means it is - # safe to use the current protocol instead, because this node was - # not changed. - key_new = (loader, PROTOCOL) - try: - node_cls = NODE_TYPE_MAPPING[key_new] - except KeyError: - # If we still cannot find the loader for this key, something - # went wrong. - type_name = f"{state['__module__']}.{state['__class__']}" - raise TypeError( - f" Can't find loader {state['__loader__']} for type {type_name} " - f"and protocol {protocol}. You might need to update skops to " - "load this file." - ) - - loaded_tree = node_cls(state, load_context, trusted=trusted) - - if allowed_types is not None and not isinstance(loaded_tree, allowed_types): - expected = " or ".join(cls.__name__ for cls in allowed_types) - raise ValueError( - f"Expected a node of type {expected}, got " - f"{type(loaded_tree).__name__}. This is probably due to a corrupted " - "or a malicious file." - ) - return loaded_tree + # What probably happened here is that we released a new protocol. If + # there is no specific key for the old protocol, it means it is safe to + # use the current protocol instead, because this node was not changed. + key_new = (loader, PROTOCOL) + try: + node_cls = NODE_TYPE_MAPPING[key_new] + except KeyError: + # If we still cannot find the loader for this key, something went + # wrong. + type_name = f"{state['__module__']}.{state['__class__']}" + raise TypeError( + f" Can't find loader {state['__loader__']} for type {type_name} and " + f"protocol {protocol}. You might need to update skops to load this " + "file." + ) + + # Check the loader before building its node, so that nothing in a rejected + # subtree is read or parsed. + _check_node_type(node_cls, allowed_types) + return node_cls(state, load_context, trusted=trusted) diff --git a/skops/io/tests/test_audit.py b/skops/io/tests/test_audit.py index d7f89869..8b758dd3 100644 --- a/skops/io/tests/test_audit.py +++ b/skops/io/tests/test_audit.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import io import json import operator @@ -406,6 +408,26 @@ def test_get_tree_allowed_types(): get_tree(state, load_context, trusted=None, allowed_types=(DictNode, ListNode)) +def test_get_tree_rejects_unexpected_node_before_building_it(): + # The loader named in the file is checked before its node is built, so a + # rejected subtree is not read at all: the unknown loader inside this list + # would raise a TypeError if the ListNode were built. + state = { + "__class__": "list", + "__module__": "builtins", + "__loader__": "ListNode", + "content": [ + {"__class__": "list", "__module__": "builtins", "__loader__": "NoSuchNode"} + ], + } + with pytest.raises(TypeError, match="Can't find loader NoSuchNode"): + get_tree(state, make_load_context(), trusted=None) + + msg = "Expected a node of type DictNode, got ListNode" + with pytest.raises(ValueError, match=msg): + get_tree(state, make_load_context(), trusted=None, allowed_types=(DictNode,)) + + @pytest.mark.parametrize( "obj, keys, expected", [