Skip to content
Open
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
Empty file added tests/__init__.py
Empty file.
76 changes: 76 additions & 0 deletions tests/test_config_utils.py
Original file line number Diff line number Diff line change
@@ -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)
45 changes: 45 additions & 0 deletions tribev2/config_utils.py
Original file line number Diff line number Diff line change
@@ -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))
10 changes: 5 additions & 5 deletions tribev2/demo_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -389,4 +389,4 @@ def predict(
n_samples,
100.0 * n_kept / max(n_samples, 1),
)
return preds, all_segments
return preds, all_segments