Skip to content

Commit 0211cce

Browse files
committed
Reject symlinked memmap collection paths
1 parent 544a4d0 commit 0211cce

3 files changed

Lines changed: 30 additions & 4 deletions

File tree

docs/source/saving.rst

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -59,9 +59,11 @@ predictable names in a shared ``/tmp`` directory.
5959

6060
Robust key encoding is enabled by default and keeps TensorDict keys within
6161
their save prefix. ``robust_key=False`` exists only to interoperate with
62-
legacy layouts; do not use it with untrusted keys or metadata. The high-level
63-
save methods overwrite existing files by default (``existsok=True``), so pass
64-
``existsok=False`` when replacement is not intended.
62+
legacy layouts; do not use it with untrusted keys or metadata. Individual
63+
memory-map leaf files overwrite existing regular files by default
64+
(``existsok=True``); pass ``existsok=False`` to reject a colliding leaf path.
65+
The metadata file is still refreshed when reusing an existing save directory,
66+
so ``existsok=False`` does not reserve the directory as a whole.
6567

6668
tensordict's memory-mapped API relies on four core methods:
6769
:meth:`~tensordict.TensorDictBase.memmap_`, :meth:`~tensordict.TensorDictBase.memmap`,

tensordict/_td.py

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3050,6 +3050,11 @@ def _memmap_(
30503050
key, robust=effective_robust_key
30513051
)
30523052
value_prefix = prefix / safe_key
3053+
if value_prefix.is_symlink():
3054+
raise RuntimeError(
3055+
"Refusing to write a memory-mapped TensorDict through "
3056+
f"symlink {value_prefix}."
3057+
)
30533058
else:
30543059
value_prefix = None
30553060
dest._tensordict[key] = value._memmap_(
@@ -3281,7 +3286,13 @@ def _make_memmap_subtd(self, key, *, robust_key):
32813286
safe_key = _encode_key_for_filesystem(
32823287
key_str, robust=effective_robust_key
32833288
)
3284-
result_tmp.memmap_(prefix=result._memmap_prefix / safe_key)
3289+
subtd_prefix = result._memmap_prefix / safe_key
3290+
if subtd_prefix.is_symlink():
3291+
raise RuntimeError(
3292+
"Refusing to write a memory-mapped TensorDict through "
3293+
f"symlink {subtd_prefix}."
3294+
)
3295+
result_tmp.memmap_(prefix=subtd_prefix)
32853296
metadata = _load_metadata(result._memmap_prefix)
32863297
_update_metadata(
32873298
metadata=metadata,

test/tensordict/test_misc.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1252,6 +1252,19 @@ def test_memmap_robust_key_nested_pathlike(self, tmp_path):
12521252
loaded_subtree = TensorDict.load_memmap(prefix, subpath=(key,))
12531253
assert_allclose_td(td[key], loaded_subtree)
12541254

1255+
def test_memmap_nested_collection_symlink_is_rejected(self, tmp_path):
1256+
prefix = tmp_path / "saved"
1257+
outside = tmp_path / "outside"
1258+
prefix.mkdir()
1259+
outside.mkdir()
1260+
(prefix / "nested").symlink_to(outside, target_is_directory=True)
1261+
td = TensorDict({"nested": {"x": torch.randn(3)}}, batch_size=[])
1262+
1263+
with pytest.raises(RuntimeError, match="TensorDict through symlink"):
1264+
td.memmap(prefix)
1265+
1266+
assert not list(outside.iterdir())
1267+
12551268
def test_memmap_robust_load_does_not_traverse_legacy_path(self, tmp_path):
12561269
key = "../outside"
12571270
td = TensorDict({key: torch.randn(3)}, batch_size=[])

0 commit comments

Comments
 (0)