Skip to content

Commit 0c2618f

Browse files
committed
Fix unsafe YAML config loading
1 parent 244df80 commit 0c2618f

4 files changed

Lines changed: 59 additions & 6 deletions

File tree

numbast/src/numbast/experimental/mlir/tools/static_binding_generator.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -45,7 +45,7 @@
4545
VERBOSE = True
4646

4747
# Register custom YAML constructor for !join tag
48-
yaml.add_constructor("!numbast_join", string_constructor)
48+
yaml.SafeLoader.add_constructor("!numbast_join", string_constructor)
4949

5050

5151
def _config_dict_uses_mlir_backend(config_dict: dict) -> bool:
@@ -59,7 +59,7 @@ def _config_dict_uses_mlir_backend(config_dict: dict) -> bool:
5959

6060
def _cfg_path_uses_mlir_backend(cfg_path: str) -> bool:
6161
with open(cfg_path) as f:
62-
config_dict = yaml.load(f, yaml.Loader)
62+
config_dict = yaml.safe_load(f)
6363
return _config_dict_uses_mlir_backend(config_dict)
6464

6565

@@ -262,7 +262,7 @@ def from_yaml_path(cls, cfg_path: str) -> "Config":
262262
A new Config instance.
263263
"""
264264
with open(cfg_path) as f:
265-
config_dict = yaml.load(f, yaml.Loader)
265+
config_dict = yaml.safe_load(f)
266266
return cls(config_dict)
267267

268268
@classmethod
Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,25 @@
1+
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2+
# SPDX-License-Identifier: Apache-2.0
3+
4+
import pytest
5+
import yaml
6+
7+
from numbast.experimental.mlir.tools import static_binding_generator as sbg
8+
9+
10+
@pytest.mark.parametrize(
11+
"load_config",
12+
[sbg._cfg_path_uses_mlir_backend, sbg.Config.from_yaml_path],
13+
)
14+
def test_config_load_rejects_python_object_tags(tmp_path, load_config):
15+
marker_path = tmp_path / "unsafe-loader-marker"
16+
cfg_path = tmp_path / "config.yaml"
17+
cfg_path.write_text(
18+
f"!!python/object/apply:builtins.open\n- {marker_path}\n- w\n",
19+
encoding="utf-8",
20+
)
21+
22+
with pytest.raises(yaml.constructor.ConstructorError):
23+
load_config(cfg_path)
24+
25+
assert not marker_path.exists()

numbast/src/numbast/tools/static_binding_generator.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,7 @@
5252
)
5353

5454
# Register custom YAML constructor for !join tag
55-
yaml.add_constructor("!numbast_join", string_constructor)
55+
yaml.SafeLoader.add_constructor("!numbast_join", string_constructor)
5656

5757

5858
def _config_dict_uses_mlir_backend(config_dict: dict) -> bool:
@@ -100,7 +100,7 @@ def _validate_mlir_backend_only_config(config_dict: dict):
100100

101101
def _cfg_path_uses_mlir_backend(cfg_path: str) -> bool:
102102
with open(cfg_path) as f:
103-
config_dict = yaml.load(f, yaml.Loader)
103+
config_dict = yaml.safe_load(f)
104104
return _config_dict_uses_mlir_backend(config_dict)
105105

106106

@@ -227,7 +227,7 @@ def from_yaml_path(cls, cfg_path: str) -> "Config":
227227
A new Config instance.
228228
"""
229229
with open(cfg_path) as f:
230-
config_dict = yaml.load(f, yaml.Loader)
230+
config_dict = yaml.safe_load(f)
231231
return cls(config_dict)
232232

233233
@classmethod

numbast/src/numbast/tools/tests/test_mlir_backend_routing.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,34 @@ def test_cfg_path_uses_mlir_backend(tmp_path):
2626
assert sbg._cfg_path_uses_mlir_backend(cfg_path)
2727

2828

29+
def test_cfg_path_supports_numbast_join_tag(tmp_path):
30+
cfg_path = tmp_path / "config.yaml"
31+
cfg_path.write_text(
32+
'MLIR Backend: !numbast_join ["tr", "ue"]\n',
33+
encoding="utf-8",
34+
)
35+
36+
assert sbg._cfg_path_uses_mlir_backend(cfg_path)
37+
38+
39+
@pytest.mark.parametrize(
40+
"load_config",
41+
[sbg._cfg_path_uses_mlir_backend, sbg.Config.from_yaml_path],
42+
)
43+
def test_config_load_rejects_python_object_tags(tmp_path, load_config):
44+
marker_path = tmp_path / "unsafe-loader-marker"
45+
cfg_path = tmp_path / "config.yaml"
46+
cfg_path.write_text(
47+
f"!!python/object/apply:builtins.open\n- {marker_path}\n- w\n",
48+
encoding="utf-8",
49+
)
50+
51+
with pytest.raises(yaml.constructor.ConstructorError):
52+
load_config(cfg_path)
53+
54+
assert not marker_path.exists()
55+
56+
2957
def test_static_generator_dispatches_mlir_backend(monkeypatch, tmp_path):
3058
class DummyConfig:
3159
mlir_backend = True

0 commit comments

Comments
 (0)