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
60 changes: 35 additions & 25 deletions src/ape/utils/os.py
Original file line number Diff line number Diff line change
Expand Up @@ -322,6 +322,15 @@ def get_package_path(package_name: str) -> Path:
return package_path


def _safe_extract_path(destination: Path, member_name: Path | str) -> Path:
"""Resolve ``destination / member_name`` and reject paths that escape ``destination``."""
dest_root = destination.resolve()
target = (destination / Path(member_name)).resolve()
if not target.is_relative_to(dest_root):
raise ValueError(f"Archive member {str(member_name)!r} escapes extract destination")
return target


def extract_archive(archive_file: Path, destination: Path | None = None):
"""
Extract an archive file. Supports ``.zip`` or ``.tar.gz``.
Expand All @@ -332,39 +341,40 @@ def extract_archive(archive_file: Path, destination: Path | None = None):
Defaults to the parent directory of the archive file.
"""
destination = destination or archive_file.parent
destination.mkdir(parents=True, exist_ok=True)
if archive_file.suffix == ".zip":
with zipfile.ZipFile(archive_file, "r") as zip_ref:
zip_members = zip_ref.namelist()
if top_level_dir := os.path.commonpath(zip_members):
for zip_member in zip_members:
# Modify the member name to remove the top-level directory.
member_path = Path(zip_member)
relative_path = (
member_path.relative_to(top_level_dir) if top_level_dir else member_path
)
target_path = destination / relative_path

if member_path.is_dir():
target_path.mkdir(parents=True, exist_ok=True)
else:
target_path.parent.mkdir(parents=True, exist_ok=True)
with zip_ref.open(member_path.as_posix()) as source:
target_path.write_bytes(source.read())

else:
zip_ref.extractall(f"{destination}")
if not zip_members:
return
top_level_dir = os.path.commonpath(zip_members)
for zip_member in zip_members:
# Modify the member name to remove the top-level directory.
member_path = Path(zip_member)
relative_path = (
member_path.relative_to(top_level_dir) if top_level_dir else member_path
)
target_path = _safe_extract_path(destination, relative_path)

if zip_member.endswith("/") or member_path.is_dir():
target_path.mkdir(parents=True, exist_ok=True)
else:
target_path.parent.mkdir(parents=True, exist_ok=True)
with zip_ref.open(zip_member) as source:
target_path.write_bytes(source.read())

elif archive_file.name.endswith(".tar.gz"):
with tarfile.open(archive_file, "r:gz") as tar_ref:
tar_members = tar_ref.getmembers()
if top_level_dir := os.path.commonpath([m.name for m in tar_members]):
for tar_member in tar_members:
# Modify the member name to remove the top-level directory.
if not tar_members:
return
top_level_dir = os.path.commonpath([m.name for m in tar_members])
for tar_member in tar_members:
# Modify the member name to remove the top-level directory.
if top_level_dir:
tar_member.name = os.path.relpath(tar_member.name, top_level_dir)
tar_ref.extract(tar_member, path=destination)

else:
tar_ref.extractall(path=f"{destination}")
_safe_extract_path(destination, tar_member.name)
tar_ref.extract(tar_member, path=destination)

else:
raise ValueError(f"Unsupported zip format: '{archive_file.suffix}'.")
Expand Down
55 changes: 55 additions & 0 deletions tests/functional/utils/test_os.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
import io
import tarfile
import zipfile
from pathlib import Path

import pytest
Expand All @@ -6,6 +9,7 @@
from ape.utils.os import (
clean_path,
create_tempdir,
extract_archive,
get_all_files_in_directory,
get_full_extension,
get_relative_path,
Expand Down Expand Up @@ -210,3 +214,54 @@ def test_clean_path_not_relative_to_home():
path = Path(name)
actual = clean_path(path)
assert actual == name


def test_extract_archive_zip_strips_top_level_dir():
with create_tempdir() as path:
archive = path / "pkg.zip"
dest = path / "out"
dest.mkdir()
with zipfile.ZipFile(archive, "w") as zf:
zf.writestr("pkg/readme.txt", "ok")
zf.writestr("pkg/src/__init__.py", "")

extract_archive(archive, dest)
assert (dest / "readme.txt").read_text() == "ok"
assert (dest / "src" / "__init__.py").read_text() == ""


def test_extract_archive_zip_rejects_parent_escape():
with create_tempdir() as path:
archive = path / "pkg.zip"
dest = path / "out"
dest.mkdir()
escaped = dest.parent / "escaped.txt"
with zipfile.ZipFile(archive, "w") as zf:
zf.writestr("pkg/readme.txt", "ok")
zf.writestr("pkg/../escaped.txt", "pwned")

with pytest.raises(ValueError, match="escapes extract destination"):
extract_archive(archive, dest)

assert not escaped.exists()


def test_extract_archive_tar_rejects_parent_escape():
with create_tempdir() as path:
archive = path / "pkg.tar.gz"
dest = path / "out"
dest.mkdir()
escaped = dest.parent / "escaped.txt"
with tarfile.open(archive, "w:gz") as tf:
readme = path / "readme.txt"
readme.write_text("ok")
tf.add(readme, arcname="pkg/readme.txt")
payload = b"pwned"
info = tarfile.TarInfo(name="pkg/../escaped.txt")
info.size = len(payload)
tf.addfile(info, io.BytesIO(payload))

with pytest.raises(ValueError, match="escapes extract destination"):
extract_archive(archive, dest)

assert not escaped.exists()
Loading