diff --git a/docs/changes.rst b/docs/changes.rst index 350b7ea1..fc97ccc7 100644 --- a/docs/changes.rst +++ b/docs/changes.rst @@ -11,6 +11,12 @@ skops Changelog 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. :pr:`547` by `Adrin Jalali`_. - Fix a regression since v0.12.0 where saving an object whose ``__reduce__`` raises failed at dump time. ``__reduce__`` is called on every object to detect a plain constructor call, but Cython extension types with a diff --git a/skops/io/_audit.py b/skops/io/_audit.py index 92c2d317..c866f2b5 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 @@ -313,10 +315,25 @@ 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, trusted: TrustedTypes | None, + *, + allowed_types: tuple[type[Node], ...] | None = None, ) -> Node: """Get the tree of nodes. @@ -341,6 +358,16 @@ 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 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 ------- loaded_tree : Node @@ -353,7 +380,9 @@ 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) + 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 @@ -378,5 +407,7 @@ def get_tree( "file." ) - loaded_tree = node_cls(state, load_context, trusted=trusted) - return loaded_tree + # 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/_general.py b/skops/io/_general.py index 097538da..486b9b02 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, @@ -492,7 +509,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): @@ -721,7 +741,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..8b758dd3 100644 --- a/skops/io/tests/test_audit.py +++ b/skops/io/tests/test_audit.py @@ -1,15 +1,22 @@ +from __future__ import annotations + import io 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 +28,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 +352,148 @@ 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)) + + +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", + [ + ({"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)