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
22 changes: 21 additions & 1 deletion src/mcore_bridge/utils/safetensors.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,26 @@ def load_slice(self, slices):
return self.loader()[slices]


def checkpoint_path(root: str, name: str) -> str:
"""Join a checkpoint-relative name and reject paths that leave ``root``.

``os.path.join`` drops the root when ``name`` is absolute, and a weight-map
entry such as ``../../other.safetensors`` would otherwise open a file
outside the model directory.
"""
if os.path.isabs(name):
raise ValueError(f'Checkpoint file {name!r} is outside {root}')
base = os.path.realpath(root)
path = os.path.realpath(os.path.join(base, name))
try:
inside = os.path.commonpath([base, path]) == base
except ValueError:
inside = False
if not inside:
raise ValueError(f'Checkpoint file {name!r} is outside {root}')
return path


class SafetensorLazyLoader:

def __init__(self, hf_model_dir: str, peft_format: bool = False):
Expand All @@ -42,7 +62,7 @@ def __init__(self, hf_model_dir: str, peft_format: bool = False):
def _open_file(self, filename: str):
"""Open a safetensors file if not already open."""
if filename not in self._file_handles:
file_path = os.path.join(self.hf_model_dir, filename)
file_path = checkpoint_path(self.hf_model_dir, filename)
self._file_handles[filename] = safe_open(file_path, framework='pt')
return self._file_handles[filename]

Expand Down
35 changes: 35 additions & 0 deletions tests/test_checkpoint_path.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
"""Weight-map filenames cannot leave the checkpoint directory."""
import ast
import os
from pathlib import Path


def _join():
path = Path(__file__).resolve().parents[1] / 'src/mcore_bridge/utils/safetensors.py'
tree = ast.parse(path.read_text())
fn = next(node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == 'checkpoint_path')
module = ast.Module(body=[fn], type_ignores=[])
ast.fix_missing_locations(module)
namespace = {'os': os}
exec(compile(module, str(path), 'exec'), namespace)
return namespace['checkpoint_path']


def test_relative_shard_stays_inside(tmp_path):
join = _join()
shard = tmp_path / 'model-00001.safetensors'
shard.write_bytes(b'x')
assert join(str(tmp_path), 'model-00001.safetensors') == os.path.realpath(shard)


def test_parent_and_absolute_names_are_rejected(tmp_path):
join = _join()
outside = tmp_path.parent / 'outside.safetensors'
outside.write_bytes(b'x')
for name in ('../outside.safetensors', str(outside)):
try:
join(str(tmp_path), name)
except ValueError:
continue
raise AssertionError(name)