Skip to content
Merged
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
4 changes: 2 additions & 2 deletions src/autogluon/cloud/predictor/cloud_predictor.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
from ..backend.constant import SAGEMAKER
from ..endpoint.endpoint import Endpoint
from ..utils.aws_utils import resolve_cloud_output_path
from ..utils.utils import unzip_file
from ..utils.utils import safe_unpack_archive

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -786,7 +786,7 @@ def _download_predictor(self, path, save_path):
s3.download_file(predictor_bucket, predictor_key_prefix, tarball_path)
logger.log(20, "Extracting the trained model tarball")
save_path = os.path.join(save_path, "AutoGluonModels")
unzip_file(tarball_path, save_path)
safe_unpack_archive(tarball_path, save_path)
return save_path

def save(self, silent: bool = False) -> None:
Expand Down
2 changes: 1 addition & 1 deletion src/autogluon/cloud/utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
is_compressed_file,
is_image_file,
read_image_bytes_and_encode,
safe_unpack_archive,
split_pred_and_pred_proba,
unzip_file,
zipfolder,
)
49 changes: 40 additions & 9 deletions src/autogluon/cloud/utils/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import logging
import os
import shutil
import stat
import tarfile
import zipfile
from datetime import datetime, timezone
Expand Down Expand Up @@ -61,15 +62,45 @@ def is_image_file(filename):
return True


def unzip_file(tarball_path, save_path):
save_path_abs = os.path.abspath(save_path)
with tarfile.open(tarball_path) as file:
for member in file.getmembers():
member_path = os.path.abspath(os.path.join(save_path_abs, member.name))
if member_path != save_path_abs and not member_path.startswith(save_path_abs + os.sep):
raise ValueError(f"Unsafe path in tarball, refusing to extract: {member.name}")
# members validated above to stay within save_path (no absolute/`..` traversal)
file.extractall(save_path) # nosec B202
# Copied verbatim from autogluon.common.loaders.load_archive.
def safe_unpack_archive(path: str | os.PathLike, dest_dir: str | os.PathLike) -> None:
"""Safely extract a zip or tar archive, rejecting path traversal and unsafe entries.

The archive format is auto-detected; supports zip and tar (incl. gz/bz2/xz).
"""
path = str(path)
dest_dir = os.path.realpath(str(dest_dir))

if zipfile.is_zipfile(path):
with zipfile.ZipFile(path, "r") as zf:
_validate_zip(zf, dest_dir)
zf.extractall(dest_dir) # nosec B202
elif tarfile.is_tarfile(path):
with tarfile.open(path) as tf:
_validate_tar(tf, dest_dir)
tf.extractall(dest_dir) # nosec B202
else:
raise ValueError(f"Unsupported or unrecognized archive format: {path}")


def _validate_zip(zf: zipfile.ZipFile, dest_dir: str) -> None:
for member in zf.infolist():
member_path = os.path.realpath(os.path.join(dest_dir, member.filename))
if not member_path.startswith(dest_dir + os.sep) and member_path != dest_dir:
raise ValueError(f"Path traversal detected: {member.filename} would extract outside {dest_dir}")
if stat.S_ISLNK(member.external_attr >> 16):
raise ValueError(f"Archive contains symlink: {member.filename}")


def _validate_tar(tf: tarfile.TarFile, dest_dir: str) -> None:
for member in tf.getmembers():
member_path = os.path.realpath(os.path.join(dest_dir, member.name))
if not member_path.startswith(dest_dir + os.sep) and member_path != dest_dir:
raise ValueError(f"Path traversal detected: {member.name} would extract outside {dest_dir}")
if member.issym() or member.islnk():
raise ValueError(f"Archive contains link: {member.name}")
if member.isdev() or member.isfifo():
raise ValueError(f"Archive contains special file: {member.name}")


def split_pred_and_pred_proba(prediction):
Expand Down
153 changes: 153 additions & 0 deletions tests/unittests/general/test_safe_unpack_archive.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
import io
import os
import tarfile
import tempfile
import zipfile

import pytest

from autogluon.cloud.utils.utils import safe_unpack_archive


def test_normal_zip_extracts_successfully():
with tempfile.TemporaryDirectory() as tmp_dir:
zip_path = os.path.join(tmp_dir, "test.zip")
with zipfile.ZipFile(zip_path, "w") as zf:
zf.writestr("data/file.txt", "hello")
zf.writestr("data/nested/file2.txt", "world")

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
safe_unpack_archive(zip_path, extract_dir)

assert os.path.exists(os.path.join(extract_dir, "data", "file.txt"))
assert os.path.exists(os.path.join(extract_dir, "data", "nested", "file2.txt"))


def test_zip_path_traversal_blocked():
with tempfile.TemporaryDirectory() as tmp_dir:
zip_path = os.path.join(tmp_dir, "test.zip")
with zipfile.ZipFile(zip_path, "w") as zf:
zf.writestr("../../etc/passwd", "pwned")

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
with pytest.raises(ValueError, match="Path traversal detected"):
safe_unpack_archive(zip_path, extract_dir)


def test_zip_symlink_blocked():
with tempfile.TemporaryDirectory() as tmp_dir:
zip_path = os.path.join(tmp_dir, "test.zip")
with zipfile.ZipFile(zip_path, "w") as zf:
info = zipfile.ZipInfo("link")
info.external_attr = 0o120777 << 16
zf.writestr(info, "/etc")

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
with pytest.raises(ValueError, match="Archive contains symlink"):
safe_unpack_archive(zip_path, extract_dir)


def test_normal_tar_extracts_successfully():
with tempfile.TemporaryDirectory() as tmp_dir:
tar_path = os.path.join(tmp_dir, "test.tar")
with tarfile.open(tar_path, "w") as tf:
data = b"hello"
info = tarfile.TarInfo("data/file.txt")
info.size = len(data)
tf.addfile(info, io.BytesIO(data))

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
safe_unpack_archive(tar_path, extract_dir)

assert os.path.exists(os.path.join(extract_dir, "data", "file.txt"))


def test_tar_path_traversal_blocked():
with tempfile.TemporaryDirectory() as tmp_dir:
tar_path = os.path.join(tmp_dir, "test.tar")
with tarfile.open(tar_path, "w") as tf:
data = b"pwned"
info = tarfile.TarInfo("../../etc/passwd")
info.size = len(data)
tf.addfile(info, io.BytesIO(data))

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
with pytest.raises(ValueError, match="Path traversal detected"):
safe_unpack_archive(tar_path, extract_dir)


def test_tar_symlink_blocked():
with tempfile.TemporaryDirectory() as tmp_dir:
tar_path = os.path.join(tmp_dir, "test.tar")
with tarfile.open(tar_path, "w") as tf:
info = tarfile.TarInfo("link")
info.type = tarfile.SYMTYPE
info.linkname = "/etc"
tf.addfile(info)

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
with pytest.raises(ValueError, match="Archive contains link"):
safe_unpack_archive(tar_path, extract_dir)


def test_tar_hardlink_blocked():
with tempfile.TemporaryDirectory() as tmp_dir:
tar_path = os.path.join(tmp_dir, "test.tar")
with tarfile.open(tar_path, "w") as tf:
info = tarfile.TarInfo("link")
info.type = tarfile.LNKTYPE
info.linkname = "/etc/passwd"
tf.addfile(info)

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
with pytest.raises(ValueError, match="Archive contains link"):
safe_unpack_archive(tar_path, extract_dir)


def test_tar_fifo_blocked():
with tempfile.TemporaryDirectory() as tmp_dir:
tar_path = os.path.join(tmp_dir, "test.tar")
with tarfile.open(tar_path, "w") as tf:
info = tarfile.TarInfo("fifo")
info.type = tarfile.FIFOTYPE
tf.addfile(info)

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
with pytest.raises(ValueError, match="Archive contains special file"):
safe_unpack_archive(tar_path, extract_dir)


def test_tar_device_blocked():
with tempfile.TemporaryDirectory() as tmp_dir:
tar_path = os.path.join(tmp_dir, "test.tar")
with tarfile.open(tar_path, "w") as tf:
info = tarfile.TarInfo("device")
info.type = tarfile.CHRTYPE
info.devmajor = 1
info.devminor = 3
tf.addfile(info)

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
with pytest.raises(ValueError, match="Archive contains special file"):
safe_unpack_archive(tar_path, extract_dir)


def test_unsupported_format_rejected():
with tempfile.TemporaryDirectory() as tmp_dir:
bad_path = os.path.join(tmp_dir, "file.txt")
with open(bad_path, "w") as file:
file.write("not an archive")

extract_dir = os.path.join(tmp_dir, "extract")
os.makedirs(extract_dir)
with pytest.raises(ValueError, match="Unsupported or unrecognized archive format"):
safe_unpack_archive(bad_path, extract_dir)
Loading