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
6 changes: 6 additions & 0 deletions docs/changes.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
37 changes: 34 additions & 3 deletions skops/io/_audit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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
Expand All @@ -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
Expand All @@ -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)
36 changes: 29 additions & 7 deletions skops/io/_general.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):
Expand Down
47 changes: 39 additions & 8 deletions skops/io/_numpy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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}.")
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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])

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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, [])
Expand Down
23 changes: 19 additions & 4 deletions skops/io/_sklearn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
6 changes: 5 additions & 1 deletion skops/io/old/_numpy_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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])
Expand Down
Loading
Loading