diff --git a/src/specify_cli/_download_security.py b/src/specify_cli/_download_security.py index 9d2d95ea72..5ff460666e 100644 --- a/src/specify_cli/_download_security.py +++ b/src/specify_cli/_download_security.py @@ -10,6 +10,7 @@ import tarfile import unicodedata import zipfile +import zlib from collections.abc import Iterator from contextlib import ExitStack, contextmanager from ipaddress import IPv4Address, IPv6Address, ip_address @@ -69,6 +70,19 @@ _BOUNDED_ZIP_COMPRESSION_METHODS = frozenset( (zipfile.ZIP_STORED, zipfile.ZIP_DEFLATED) ) +#: Decompression failures a truncated or corrupt gzip stream raises from +#: ``tarfile``. Most are wrapped in ``TarError``, but two escape raw, and +#: neither derives from ``TarError`` or ``OSError``, so both bypass a +#: ``(TarError, OSError)`` handler: +#: +#: * ``EOFError`` -- from the gzip layer when the stream ends before its +#: end-of-stream marker, i.e. a truncated archive. +#: * ``zlib.error`` -- from a corrupt deflate block. ``tarfile`` converts this +#: to ``ReadError`` while reading a member *header*, but the forward seek it +#: performs to skip member *data* sits outside that conversion, so a corrupt +#: region past the first header escapes raw. +_TAR_DECOMPRESSION_ERRORS = (tarfile.TarError, EOFError, zlib.error) + _ARCHIVE_CONTENT_TYPES: dict[str, ArchiveFormat] = { "application/gzip": "tar.gz", "application/x-gzip": "tar.gz", @@ -166,7 +180,11 @@ def detect_archive_format( try: with tarfile.open(fileobj=archive_file, mode="r:gz"): is_tar_gz = True - except tarfile.TarError: + except _TAR_DECOMPRESSION_ERRORS: + # A truncated gzip stream raises a bare EOFError here rather + # than a TarError, so catching only TarError let it escape + # this probe as a raw exception instead of leaving + # ``is_tar_gz`` False and reporting the format mismatch. pass archive_file.seek(0) except OSError as exc: @@ -1077,7 +1095,7 @@ def safe_extract_tar( mode="r:gz", fileobj=archive_file, ) - except (tarfile.TarError, OSError) as exc: + except (*_TAR_DECOMPRESSION_ERRORS, OSError) as exc: _raise_from(error_type, f"Invalid tar.gz archive: {archive_path}", exc) with archive: @@ -1149,7 +1167,7 @@ def safe_extract_tar( f"of {max_total_bytes} bytes", ) validated.append((member, normalized_name, is_dir)) - except (tarfile.TarError, OSError) as exc: + except (*_TAR_DECOMPRESSION_ERRORS, OSError) as exc: _raise_from( error_type, f"Invalid tar.gz archive: {archive_path}", diff --git a/tests/test_download_security.py b/tests/test_download_security.py index df6f9180d4..39a566e48d 100644 --- a/tests/test_download_security.py +++ b/tests/test_download_security.py @@ -475,6 +475,165 @@ def test_safe_extract_tar_enforces_entry_and_size_limits(tmp_path): safe_extract_tar(archive_path, tmp_path / "total", max_total_bytes=7) +def _truncated_tar_gz_bytes(keep_bytes): + """Return the leading *keep_bytes* of a multi-member tar.gz's bytes. + + A gzip stream cut short this way ends before its end-of-stream marker, so + reading it raises a bare ``EOFError`` from the gzip layer. ``tarfile`` + decompresses lazily, so *where* that surfaces depends on how much is kept: + a very short prefix fails in ``tarfile.open`` itself, while a longer one + opens fine and only fails once members are iterated. + """ + buffer = io.BytesIO() + with tarfile.open(fileobj=buffer, mode="w:gz") as archive: + for index in range(5): + info = tarfile.TarInfo(f"file{index}.txt") + content = bytes(range(256)) * 400 + info.size = len(content) + archive.addfile(info, io.BytesIO(content)) + return buffer.getvalue()[:keep_bytes] + + +def test_detect_archive_format_rejects_truncated_tar_gz(tmp_path): + # A gzip stream truncated before tarfile can read its first header raises a + # bare EOFError -- not a TarError -- from the format probe. Catching only + # TarError let it escape as a raw exception instead of leaving is_tar_gz + # False and reporting the module's clean format-mismatch error. + archive_path = tmp_path / "truncated.tar.gz" + archive_path.write_bytes(_truncated_tar_gz_bytes(64)) + + with pytest.raises(ValueError, match="format mismatch"): + detect_archive_format(archive_path) + + +@pytest.mark.parametrize("keep_bytes", [64, 512, 2048]) +def test_safe_extract_tar_rejects_truncated_archive(tmp_path, keep_bytes): + # The same bare EOFError, from tarfile.open on a short prefix and from + # member iteration on a longer one. Both sites reported it raw. + archive_path = tmp_path / f"truncated-{keep_bytes}.tar.gz" + archive_path.write_bytes(_truncated_tar_gz_bytes(keep_bytes)) + + with pytest.raises(ValueError, match="Invalid tar.gz archive"): + safe_extract_tar(archive_path, tmp_path / f"out-{keep_bytes}") + + +def test_safe_extract_tar_wraps_truncation_in_caller_error_type(tmp_path): + # The leak bypassed the caller's domain error type entirely, so callers + # that only catch their own error (or ValueError) crashed the command. + archive_path = tmp_path / "truncated.tar.gz" + archive_path.write_bytes(_truncated_tar_gz_bytes(2048)) + + with pytest.raises(_CustomZipError, match="Invalid tar.gz archive"): + safe_extract_tar( + archive_path, + tmp_path / "out", + error_type=_CustomZipError, + ) + + +def test_safe_extract_archive_rejects_truncated_tar_gz(tmp_path): + archive_path = tmp_path / "truncated.tar.gz" + archive_path.write_bytes(_truncated_tar_gz_bytes(2048)) + + with pytest.raises(ValueError): + safe_extract_archive(archive_path, tmp_path / "out") + + +def _corrupt_deflate_tar_gz_bytes(): + """Return a tar.gz whose deflate stream is corrupt mid-member. + + Unlike truncation, which the gzip layer reports as ``EOFError``, mangling + bytes inside a deflate block raises ``zlib.error``. ``tarfile`` converts + that to ``ReadError`` when it surfaces while reading a member *header*, but + the forward seek it performs to skip over member *data* sits outside that + conversion, so the raw ``zlib.error`` escapes from there. + + Reaching that seek requires members larger than the gzip read buffer -- + with small members the whole stream is decompressed during the first header + read, and the error is wrapped. Hence two 256 KiB members, stored at + ``compresslevel=1`` so the fixture stays a few kilobytes on disk, with the + corruption placed past the midpoint so the first header still reads clean. + """ + buffer = io.BytesIO() + with tarfile.open(fileobj=buffer, mode="w:gz", compresslevel=1) as archive: + for index in range(2): + info = tarfile.TarInfo(f"file{index}.txt") + content = bytes((i * 7 + index) % 256 for i in range(1024)) * 256 + info.size = len(content) + archive.addfile(info, io.BytesIO(content)) + + raw = bytearray(buffer.getvalue()) + midpoint = len(raw) // 2 + for offset in range(midpoint, min(midpoint + 64, len(raw) - 8)): + raw[offset] ^= 0xA5 + return bytes(raw) + + +def test_corrupt_deflate_fixture_raises_bare_zlib_error(): + # Guards the fixture itself: the tests below are only meaningful while this + # archive reaches the module as a bare zlib.error -- neither a TarError nor + # an OSError, so a (TarError, OSError) handler would miss it. If a future + # Python wraps it, this fails loudly instead of the coverage silently + # decaying into a duplicate of the EOFError cases. + archive_file = io.BytesIO(_corrupt_deflate_tar_gz_bytes()) + + with tarfile.open(fileobj=archive_file, mode="r:gz") as archive: + with pytest.raises(zlib.error): + for _member in archive: + pass + + +def test_detect_archive_format_accepts_corrupt_deflate_tar_gz(tmp_path): + # Detection is a format probe, not an integrity check: tarfile.open reads + # only the first member header, which is intact here, so the archive is + # correctly identified as tar.gz and the corruption is caught later by + # safe_extract_tar (see the tests below). + # + # Note this does not exercise the probe's zlib.error handling, which is + # unreachable: the header read is inside tarfile's own + # zlib.error -> ReadError conversion, so the probe sees ReadError. The + # zlib.error arm of _TAR_DECOMPRESSION_ERRORS is defensive at this site and + # load-bearing only at the two safe_extract_tar sites. + archive_path = tmp_path / "corrupt.tar.gz" + archive_path.write_bytes(_corrupt_deflate_tar_gz_bytes()) + + assert detect_archive_format(archive_path) == "tar.gz" + + +def test_safe_extract_tar_rejects_corrupt_deflate(tmp_path): + archive_path = tmp_path / "corrupt.tar.gz" + archive_path.write_bytes(_corrupt_deflate_tar_gz_bytes()) + + with pytest.raises(ValueError, match="Invalid tar.gz archive"): + safe_extract_tar(archive_path, tmp_path / "out") + + +def test_safe_extract_tar_wraps_corrupt_deflate_in_caller_error_type(tmp_path): + # zlib.error must reach the caller's domain error type, exactly as EOFError + # does, so this cannot regress independently of the truncation handling. + archive_path = tmp_path / "corrupt.tar.gz" + archive_path.write_bytes(_corrupt_deflate_tar_gz_bytes()) + + with pytest.raises(_CustomZipError, match="Invalid tar.gz archive"): + safe_extract_tar( + archive_path, + tmp_path / "out", + error_type=_CustomZipError, + ) + + +def test_safe_extract_archive_wraps_corrupt_deflate_in_caller_error_type(tmp_path): + archive_path = tmp_path / "corrupt.tar.gz" + archive_path.write_bytes(_corrupt_deflate_tar_gz_bytes()) + + with pytest.raises(_CustomZipError, match="Invalid tar.gz archive"): + safe_extract_archive( + archive_path, + tmp_path / "out", + error_type=_CustomZipError, + ) + + @pytest.mark.parametrize("suffix", [".zip", ".tar.gz", ".tgz"]) def test_safe_extract_archive_has_format_parity(tmp_path, suffix): archive_path = tmp_path / f"package{suffix}"