diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/test_config_utils.py b/tests/test_config_utils.py new file mode 100644 index 00000000..0904becc --- /dev/null +++ b/tests/test_config_utils.py @@ -0,0 +1,76 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import textwrap + +import pytest +import yaml + +from tribev2.config_utils import TribeConfigLoader, load_config + + +def test_load_config_accepts_python_tuple_tag(tmp_path): + """The published facebook/tribev2 config uses !!python/tuple; we must + keep loading it as an actual tuple to stay backward compatible.""" + config_path = tmp_path / "config.yaml" + config_path.write_text( + textwrap.dedent("""\ + workdir: + excludes: !!python/tuple + - __pycache__ + - .git + data: + study: + path: . + """), + encoding="utf-8", + ) + + config = load_config(config_path) + + assert isinstance(config["workdir"]["excludes"], tuple) + assert config["workdir"]["excludes"] == ("__pycache__", ".git") + assert config["data"]["study"]["path"] == "." + + +def test_load_config_handles_plain_yaml(tmp_path): + """Configs without any Python-specific tags must still load cleanly.""" + config_path = tmp_path / "config.yaml" + config_path.write_text( + textwrap.dedent("""\ + data: + study: + path: . + infra: + folder: /tmp/out + average_subjects: true + """), + encoding="utf-8", + ) + + config = load_config(config_path) + + assert config["data"]["study"]["path"] == "." + assert config["infra"]["folder"] == "/tmp/out" + assert config["average_subjects"] is True + + +def test_loader_rejects_python_object_new_tags(): + """The whole point of the change: arbitrary Python object construction + via !!python/object/new:* must be rejected.""" + payload = 'cmd: !!python/object/new:os.system ["echo hacked"]\n' + + with pytest.raises(yaml.constructor.ConstructorError): + yaml.load(payload, Loader=TribeConfigLoader) + + +def test_loader_rejects_python_name_tags(): + """!!python/name lookups (another arbitrary-import vector) must also + be rejected.""" + payload = "cmd: !!python/name:os.system\n" + + with pytest.raises(yaml.constructor.ConstructorError): + yaml.load(payload, Loader=TribeConfigLoader) diff --git a/tribev2/config_utils.py b/tribev2/config_utils.py new file mode 100644 index 00000000..3b4c3ce7 --- /dev/null +++ b/tribev2/config_utils.py @@ -0,0 +1,45 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +"""Safer YAML loading for TRIBE config files. + +The pretrained-model loader used to call ``yaml.load`` with +``yaml.UnsafeLoader``, which constructs arbitrary Python objects from +YAML tags. That allows ``!!python/object/new:os.system [...]``-style +payloads to execute code at load time. + +This module exposes ``load_config`` which uses a ``yaml.SafeLoader`` +subclass that explicitly opts in to ``!!python/tuple``, the only +Python-specific tag known to appear in published TRIBE config files. +All other Python-specific tags (``!!python/object``, +``!!python/object/new``, ``!!python/name``, ``!!python/module``, …) are +rejected by the safe base class. +""" + +from pathlib import Path + +import yaml +from exca import ConfDict + + +class TribeConfigLoader(yaml.SafeLoader): + """Safe YAML loader for TRIBE configs with explicit tuple support.""" + + +def _construct_python_tuple(loader: yaml.Loader, node: yaml.Node) -> tuple: + return tuple(loader.construct_sequence(node)) + + +TribeConfigLoader.add_constructor( + "tag:yaml.org,2002:python/tuple", + _construct_python_tuple, +) + + +def load_config(path: str | Path) -> ConfDict: + """Load a TRIBE config YAML safely and return a ConfDict.""" + with open(path, "r", encoding="utf-8") as f: + return ConfDict(yaml.load(f, Loader=TribeConfigLoader)) diff --git a/tribev2/demo_utils.py b/tribev2/demo_utils.py index fa735a79..32002c49 100644 --- a/tribev2/demo_utils.py +++ b/tribev2/demo_utils.py @@ -15,11 +15,12 @@ import pydantic import requests import torch -import yaml from einops import rearrange -from exca import ConfDict, TaskInfra +from exca import TaskInfra from tqdm import tqdm +from tribev2.config_utils import load_config + logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) if not logger.handlers: @@ -201,8 +202,7 @@ def from_pretrained( repo_id = str(checkpoint_dir) config_path = hf_hub_download(repo_id, "config.yaml") ckpt_path = hf_hub_download(repo_id, checkpoint_name) - with open(config_path, "r") as f: - config = ConfDict(yaml.load(f, Loader=yaml.UnsafeLoader)) + config = load_config(config_path) for modality in ["text", "audio", "video"]: config[f"data.{modality}_feature.infra.folder"] = cache_folder config[f"data.{modality}_feature.infra.cluster"] = cluster @@ -389,4 +389,4 @@ def predict( n_samples, 100.0 * n_kept / max(n_samples, 1), ) - return preds, all_segments + return preds, all_segments \ No newline at end of file