From c2c2f3e5450bfd8970d2fd9b69fa39b51ebfad37 Mon Sep 17 00:00:00 2001 From: Julien Cristau Date: Fri, 25 Sep 2026 12:20:18 +0000 Subject: [PATCH 1/4] perf(fetch-content): let tar decompress zstd archives When extracting a zstd compressed tarball on a non-Windows platform and a zstd program is available, run `tar --use-compress-program=zstd -xf` instead of decompressing in Python and piping into `tar xf -`. This moves decompression out of the fetch-content process, where it competes for the GIL with concurrent downloads. Extracting a 1.5GB toolchain from local disk takes 7.3-8.6s this way, against 10.5-12s for the pipe. The long option is used because GNU tar before 1.31 has no --zstd, and bsdtar on macOS accepts it too. Without a zstd program, the existing pipe is used. --- src/taskgraph/run-task/fetch-content | 12 ++++ test/test_scripts_fetch_content.py | 93 ++++++++++++++++++++++++++++ 2 files changed, 105 insertions(+) diff --git a/src/taskgraph/run-task/fetch-content b/src/taskgraph/run-task/fetch-content index 337d807fb..f189d2df7 100755 --- a/src/taskgraph/run-task/fetch-content +++ b/src/taskgraph/run-task/fetch-content @@ -368,6 +368,11 @@ def open_stream(path: pathlib.Path): raise ArchiveTypeNotSupported(path) +def is_zstd(path: pathlib.Path): + with path.open(mode="rb") as fh: + return fh.read(4) == b"\x28\xb5\x2f\xfd" + + def archive_type(path: pathlib.Path): """Attempt to identify a path as an extractable archive.""" if path.suffixes[-2:-1] == [".tar"] or path.suffixes[-1:] == [".tgz"]: @@ -466,6 +471,13 @@ def extract_archive(path, dest_dir): tar = TarFile.open(fileobj=ifh, mode="r|") tar.extractall(str(dest_dir)) args = [] + elif is_zstd(path) and shutil.which("zstd"): + # Decompressing in a separate process keeps it from competing + # with concurrent downloads for the GIL. + ifh.close() + ifh = open(os.devnull, "rb") + args = ["tar", "--use-compress-program=zstd", "-xf", str(path)] + pipe_stdin = False else: args = ["tar", "xf", "-"] pipe_stdin = True diff --git a/test/test_scripts_fetch_content.py b/test/test_scripts_fetch_content.py index a971c22c5..e3320da75 100644 --- a/test/test_scripts_fetch_content.py +++ b/test/test_scripts_fetch_content.py @@ -1,5 +1,6 @@ import io import json +import lzma import os import pathlib import shutil @@ -379,3 +380,95 @@ def test_merge_tree_readonly_dir_from_later_fetch(tmp_path, fetch_content_mod): assert (dest / "tests" / "a.txt").read_text() == "a" assert (dest / "tests" / "b.txt").read_text() == "b" assert stat.S_IMODE((dest / "tests").stat().st_mode) == 0o555 + + +@pytest.fixture +def popen_calls(monkeypatch, fetch_content_mod): + """Record the arguments of every subprocess.Popen call, letting them run.""" + calls = [] + real_popen = fetch_content_mod.subprocess.Popen + + def recording_popen(args, *a, **kw): + calls.append(args) + return real_popen(args, *a, **kw) + + monkeypatch.setattr(fetch_content_mod.subprocess, "Popen", recording_popen) + return calls + + +def _tar_bytes(tmp_path, files): + _make_tar(tmp_path / "raw.tar", files) + return (tmp_path / "raw.tar").read_bytes() + + +def _make_tar_zst(path, files): + zstandard = pytest.importorskip("zstandard") + path.write_bytes( + zstandard.ZstdCompressor().compress(_tar_bytes(path.parent, files)) + ) + + +@pytest.mark.skipif(sys.platform == "win32", reason="Windows extracts with tarfile") +@pytest.mark.skipif(not shutil.which("zstd"), reason="needs the zstd program") +def test_extract_archive_zstd_with_tar(tmp_path, fetch_content_mod, popen_calls): + archive = tmp_path / "archive.tar.zst" + _make_tar_zst(archive, {"dir/a.txt": "a", "b.txt": "b"}) + dest = tmp_path / "dest" + dest.mkdir() + + fetch_content_mod.extract_archive(archive, dest) + + assert popen_calls == [ + ["tar", "--use-compress-program=zstd", "-xf", str(archive.resolve())] + ] + assert (dest / "dir" / "a.txt").read_text() == "a" + assert (dest / "b.txt").read_text() == "b" + + +@pytest.mark.skipif(sys.platform == "win32", reason="Windows extracts with tarfile") +def test_extract_archive_zstd_without_zstd_program( + tmp_path, fetch_content_mod, popen_calls, monkeypatch +): + """Without a zstd program, decompress in Python and pipe to tar.""" + archive = tmp_path / "archive.tar.zst" + _make_tar_zst(archive, {"dir/a.txt": "a"}) + dest = tmp_path / "dest" + dest.mkdir() + monkeypatch.setattr(fetch_content_mod.shutil, "which", lambda name: None) + + fetch_content_mod.extract_archive(archive, dest) + + assert popen_calls == [["tar", "xf", "-"]] + assert (dest / "dir" / "a.txt").read_text() == "a" + + +@pytest.mark.skipif(sys.platform == "win32", reason="Windows extracts with tarfile") +def test_extract_archive_other_compression_pipes_to_tar( + tmp_path, fetch_content_mod, popen_calls +): + """Only zstd is handed to tar; other formats still go through the pipe.""" + archive = tmp_path / "archive.tar.xz" + archive.write_bytes(lzma.compress(_tar_bytes(tmp_path, {"dir/a.txt": "a"}))) + dest = tmp_path / "dest" + dest.mkdir() + + fetch_content_mod.extract_archive(archive, dest) + + assert popen_calls == [["tar", "xf", "-"]] + assert (dest / "dir" / "a.txt").read_text() == "a" + + +@pytest.mark.skipif(sys.platform == "win32", reason="Windows extracts with tarfile") +@pytest.mark.skipif(not shutil.which("zstd"), reason="needs the zstd program") +def test_extract_archive_zstd_tar_failure(tmp_path, fetch_content_mod): + """A corrupt archive makes tar exit non-zero, which must be reported.""" + archive = tmp_path / "archive.tar.zst" + # Incompressible, so that the tar header survives the truncation. + _make_tar_zst(archive, {"a.txt": os.urandom(1000000).hex()}) + data = archive.read_bytes() + archive.write_bytes(data[: len(data) // 2]) + dest = tmp_path / "dest" + dest.mkdir() + + with pytest.raises(Exception, match="exited"): + fetch_content_mod.extract_archive(archive, dest) From c313bfe85acaed3bec6cad73f3d85abe3413e583 Mon Sep 17 00:00:00 2001 From: Julien Cristau Date: Fri, 25 Sep 2026 12:20:30 +0000 Subject: [PATCH 2/4] perf(fetch-content): read downloads in 1MiB chunks stream_download read responses 64KiB at a time. Add a CHUNK_SIZE of 1MiB and use it instead. In a local benchmark of Python's HTTPS read loop with 8 threads, 1MiB reads give about 30% more throughput than 64KiB ones. --- src/taskgraph/run-task/fetch-content | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/taskgraph/run-task/fetch-content b/src/taskgraph/run-task/fetch-content index f189d2df7..ba842b7df 100755 --- a/src/taskgraph/run-task/fetch-content +++ b/src/taskgraph/run-task/fetch-content @@ -43,6 +43,8 @@ except ImportError: CONCURRENCY = multiprocessing.cpu_count() +CHUNK_SIZE = 1024 * 1024 + def log(msg): print(msg, file=sys.stderr) @@ -219,7 +221,7 @@ def stream_download(url, sha256=None, size=None, headers=None): size = content_length while True: - chunk = fh.read(65536) + chunk = fh.read(CHUNK_SIZE) if not chunk: break From 41ff7f0c87871b9ba7954b605e4dfcc84d596f5e Mon Sep 17 00:00:00 2001 From: Julien Cristau Date: Fri, 25 Sep 2026 16:42:41 +0000 Subject: [PATCH 3/4] refactor(fetch-content): detect archives from file objects, share the tar pipe Split open_stream into open_fileobj, which works on any buffered file object and peeks at the magic instead of seeking, and move the loop feeding an extraction program its input into run_pipeline, which can also chain several programs. A process that exits early no longer turns into a BrokenPipeError hiding its exit status. Both are needed to extract archives while they download. --- src/taskgraph/run-task/fetch-content | 72 +++++++++++++++++++++------- 1 file changed, 54 insertions(+), 18 deletions(-) diff --git a/src/taskgraph/run-task/fetch-content b/src/taskgraph/run-task/fetch-content index ba842b7df..fc1bd7d49 100755 --- a/src/taskgraph/run-task/fetch-content +++ b/src/taskgraph/run-task/fetch-content @@ -8,6 +8,7 @@ import bz2 import concurrent.futures import contextlib import datetime +import functools import gzip import hashlib import io @@ -338,12 +339,19 @@ class ArchiveTypeNotSupported(Exception): super(Exception, self).__init__(f"Archive type not supported for {path}") +ZSTD_MAGIC = b"\x28\xb5\x2f\xfd" + + def open_stream(path: pathlib.Path): """Attempt to identify a path as an extractable archive by looking at its content.""" - fh = path.open(mode="rb") - magic = fh.read(6) - fh.seek(0) + return open_fileobj(path.open(mode="rb"), path) + + +def open_fileobj(fh, name): + """Like ``open_stream``, for a buffered file object that may not be + seekable.""" + magic = fh.peek(6)[:6] if magic[:2] == b"PK": return "zip", fh if magic[:2] == b"\x1f\x8b": @@ -352,7 +360,7 @@ def open_stream(path: pathlib.Path): fh = bz2.BZ2File(fh) elif magic == b"\xfd7zXZ\x00": fh = lzma.LZMAFile(fh) - elif magic[:4] == b"\x28\xb5\x2f\xfd": + elif magic[:4] == ZSTD_MAGIC: fh = ZstdDecompressor().stream_reader(fh) fh = io.BufferedReader(fh) try: @@ -367,12 +375,12 @@ def open_stream(path: pathlib.Path): return "tar", fh except Exception: pass - raise ArchiveTypeNotSupported(path) + raise ArchiveTypeNotSupported(name) def is_zstd(path: pathlib.Path): with path.open(mode="rb") as fh: - return fh.read(4) == b"\x28\xb5\x2f\xfd" + return fh.read(4) == ZSTD_MAGIC def archive_type(path: pathlib.Path): @@ -492,24 +500,52 @@ def extract_archive(path, dest_dir): raise ValueError(f"unknown archive format: {path}") if args: - with ifh, subprocess.Popen( - args, cwd=str(dest_dir), bufsize=0, stdin=subprocess.PIPE - ) as p: - while True: - if not pipe_stdin: - break + with ifh: + chunks = iter(functools.partial(ifh.read, 131072), b"") + run_pipeline([args], dest_dir, chunks if pipe_stdin else ()) + + log(f"{path} extracted in {time.time() - t0:.3f}s") + - chunk = ifh.read(131072) - if not chunk: - break +def run_pipeline(commands, cwd, chunks): + """Run ``commands`` in ``cwd``, each reading the output of the previous + one, and the first reading ``chunks``.""" + procs = [] + stdin = subprocess.PIPE + for i, args in enumerate(commands): + last = i == len(commands) - 1 + p = subprocess.Popen( + args, + cwd=str(cwd), + bufsize=0, + stdin=stdin, + stdout=None if last else subprocess.PIPE, + ) + if procs: + # Only the next process holds the pipe, so that the previous one + # gets SIGPIPE if it exits. + procs[-1].stdout.close() + stdin = p.stdout + procs.append(p) - p.stdin.write(chunk) + try: + for chunk in chunks: + procs[0].stdin.write(chunk) + except BrokenPipeError: + # A process exited early; its exit status says why. + pass + finally: + try: + procs[0].stdin.close() + except BrokenPipeError: + pass + for p in procs: + p.wait() + for args, p in zip(commands, procs): if p.returncode: raise Exception(f"{args!r} exited {p.returncode}") - log(f"{path} extracted in {time.time() - t0:.3f}s") - def should_repack_archive( orig: pathlib.Path, dest: pathlib.Path, strip_components=0, add_prefix="" From f010158b28495473de2ed92356e0e10f79107c64 Mon Sep 17 00:00:00 2001 From: Julien Cristau Date: Fri, 25 Sep 2026 16:48:32 +0000 Subject: [PATCH 4/4] perf(fetch-content): extract tar archives while they download Extracted archives are written to disk only to be read back and deleted. When the fetch phase is bound by disk writes, that is wasted work. For example, a Firefox build-win64/opt task on a GCP c3d-standard-16-lssd worker writes about 15GiB while fetching, past the point where the kernel throttles writers to disk speed, and 3.4GiB of that is the archives. Streaming them straight into tar, measured with curl on those workers, took that fetch phase from 18.4s to 12.0s. When a fetch is extracted, feed the response straight to the extractor instead: `zstd -dc | tar xf -` for zstd when the zstd program is available, Python decompression piped to `tar xf -` otherwise, and tarfile on Windows. zstd isn't left to `tar --use-compress-program`, because tar relays stdin to it in small blocks, which is several times slower. The start of the stream is inspected the same way as a downloaded file is, and kept, so that zip archives and files that aren't archives can still be written out whole from the same request. Zip archives, detected by name, and fetches that aren't extracted keep going through a file. Extraction goes into the fetch's staging directory, which is only merged once the download has been read to the end and its size and sha256 verified. On any failure the staging directory is emptied and the fetch starts again from scratch, unless the download itself turns out to be intact, in which case the archive is broken and the error is raised straight away rather than retried. Streaming is on by default; set TASKGRAPH_FETCH_STREAM=0 to turn it off. --- docs/howto/use-fetches.rst | 13 + src/taskgraph/run-task/fetch-content | 182 +++++++++- test/test_scripts_fetch_content.py | 519 ++++++++++++++++++++++++++- 3 files changed, 698 insertions(+), 16 deletions(-) diff --git a/docs/howto/use-fetches.rst b/docs/howto/use-fetches.rst index d7e727b83..c6996dae9 100644 --- a/docs/howto/use-fetches.rst +++ b/docs/howto/use-fetches.rst @@ -105,3 +105,16 @@ There are a few differences from the earlier ``build`` examples here: It is not possible to configure the ``dest`` or ``extract`` values when using ``fetch`` or ``toolchain`` kinds. + +Tuning Download Performance +--------------------------- + +The following environment variable tunes how ``fetch-content`` downloads +artifacts. + +``TASKGRAPH_FETCH_STREAM`` + Tar archives that are extracted are extracted while they download, rather + than written to disk and read back, into a staging directory that is only + merged into place once the download is complete and verified. This is on + by default; set to ``0`` to download to a file first. Zip archives, and + fetches that aren't extracted, always go through a file. diff --git a/src/taskgraph/run-task/fetch-content b/src/taskgraph/run-task/fetch-content index fc1bd7d49..6c8390f8c 100755 --- a/src/taskgraph/run-task/fetch-content +++ b/src/taskgraph/run-task/fetch-content @@ -12,6 +12,7 @@ import functools import gzip import hashlib import io +import itertools import json import lzma import multiprocessing @@ -44,6 +45,10 @@ except ImportError: CONCURRENCY = multiprocessing.cpu_count() +# Whether to extract tar archives while they download, rather than writing +# them to disk and reading them back. TASKGRAPH_FETCH_STREAM=0 disables it. +DEFAULT_STREAMING = True + CHUNK_SIZE = 1024 * 1024 @@ -251,6 +256,13 @@ def stream_download(url, sha256=None, size=None, headers=None): ) +def streaming_enabled(): + value = os.environ.get("TASKGRAPH_FETCH_STREAM", "").strip().lower() + if not value: + return DEFAULT_STREAMING + return value in ("1", "true", "yes") + + def download_to_path(url, path, sha256=None, size=None, headers=None): """Download a URL to a filesystem path, possibly with verification.""" @@ -547,6 +559,166 @@ def run_pipeline(commands, cwd, chunks): raise Exception(f"{args!r} exited {p.returncode}") +class IterReader(io.RawIOBase): + """A file object reading from an iterator of chunks. + + Until ``iter_from_start`` or ``forget`` is called, it keeps what it has + read, so that the start of a stream can be inspected and then handed on + whole. + """ + + def __init__(self, chunks): + self._chunks = chunks + self._chunk = b"" + self._pos = 0 + self._recorded = [] + # Whether reading from the iterator raised. + self.failed = False + + def readable(self): + return True + + def _next(self): + try: + chunk = next(self._chunks, b"") + except Exception: + self.failed = True + raise + if self._recorded is not None and chunk: + self._recorded.append(chunk) + return chunk + + def readinto(self, b): + if self._pos == len(self._chunk): + self._chunk = memoryview(self._next()) + self._pos = 0 + n = min(len(b), len(self._chunk) - self._pos) + b[:n] = self._chunk[self._pos : self._pos + n] + self._pos += n + return n + + def _rest(self): + while True: + chunk = self._next() + if not chunk: + return + yield chunk + + def iter_from_start(self): + """Iterate over the whole stream, from the start.""" + recorded, self._recorded = self._recorded, None + return itertools.chain(recorded, self._rest()) + + def forget(self): + """Stop keeping what has been read.""" + self._recorded = None + + def drain(self): + """Read the rest of the stream, discarding it.""" + self.forget() + for _ in self._rest(): + pass + + +def extract_stream(reader, dest_dir): + """Extract the tar archive read from ``reader`` into ``dest_dir``. + + Returns False, having only read the start of the stream, if it isn't a + tar archive. Otherwise reads the stream to the end, so that the download + it comes from gets verified. + """ + fh = io.BufferedReader(reader, CHUNK_SIZE) + magic = fh.peek(4)[:4] + try: + typ, ifh = open_fileobj(fh, "stream") + except (ArchiveTypeNotSupported, ValueError): + # ValueError is a zstd archive without the zstandard module, which + # extract_archive will report. + return False + if typ != "tar": + return False + + if sys.platform == "win32": + # See extract_archive. + reader.forget() + TarFile.open(fileobj=ifh, mode="r|").extractall(str(dest_dir)) + elif magic == ZSTD_MAGIC and shutil.which("zstd"): + # Not tar --use-compress-program: when reading from stdin, tar relays + # it to the decompressor in small blocks, which is several times + # slower. + run_pipeline( + [["zstd", "-dcq"], ["tar", "xf", "-"]], + dest_dir, + reader.iter_from_start(), + ) + else: + reader.forget() + run_pipeline( + [["tar", "xf", "-"]], + dest_dir, + iter(functools.partial(ifh.read1, CHUNK_SIZE), b""), + ) + reader.drain() + return True + + +def empty_directory(path): + """Remove everything in ``path``, including read-only directories.""" + for root, dirs, _ in os.walk(path): + for d in dirs: + d = os.path.join(root, d) + if not os.path.islink(d): + os.chmod(d, stat.S_IRWXU) + for entry in os.scandir(path): + if entry.is_dir(follow_symlinks=False): + shutil.rmtree(entry.path) + else: + os.unlink(entry.path) + + +def stream_and_extract(url, dest_path, dest_dir, sha256=None, size=None, headers=None): + """Download ``url`` and extract it into ``dest_dir`` as it arrives. + + If it turns out not to be a tar archive, it is downloaded to + ``dest_path`` instead, and False is returned. + + With ``sha256`` or ``size``, content is extracted before it can be + verified. That's fine as long as ``dest_dir`` is a staging directory that + is only merged once this returns: on a mismatch it is emptied, the same + as a download that fails verification is deleted. + """ + for _ in retrier(attempts=5, sleeptime=60): + chunks = stream_download(url, sha256=sha256, size=size, headers=headers) + reader = IterReader(chunks) + try: + t0 = time.time() + if extract_stream(reader, dest_dir): + log(f"{url} extracted to {dest_dir} in {time.time() - t0:.3f}s") + return True + log(f"{url} is not a tar archive; downloading it to {dest_path}") + with rename_after_close(dest_path, "wb") as fh: + for chunk in reader.iter_from_start(): + fh.write(chunk) + return False + except Exception as e: + log(f"Download failed: {e}") + empty_directory(dest_dir) + if not reader.failed: + # The download didn't fail, so find out whether it is intact. + # If it is, the archive itself is broken and retrying won't + # help. + try: + reader.drain() + except Exception: + pass + else: + raise + finally: + chunks.close() + + raise Exception("Download failed, no more retries!") + + def should_repack_archive( orig: pathlib.Path, dest: pathlib.Path, strip_components=0, add_prefix="" ) -> bool: @@ -769,12 +941,20 @@ def fetch_and_extract(url, dest_dir, extract=True, sha256=None, size=None): If the downloaded URL is an archive, it is extracted automatically and the archive is deleted. Otherwise the file remains in place in the destination directory. + + When extracting, ``dest_dir`` must be a staging directory: a failed + attempt empties it. """ basename = urllib.parse.unquote(urllib.parse.urlparse(url).path.split("/")[-1]) dest_path = dest_dir / basename - download_to_path(url, dest_path, sha256=sha256, size=size) + # Zip archives can't be streamed, so they keep going through a file. + if extract and streaming_enabled() and archive_type(dest_path) != "zip": + if stream_and_extract(url, dest_path, dest_dir, sha256=sha256, size=size): + return + else: + download_to_path(url, dest_path, sha256=sha256, size=size) if not extract: return diff --git a/test/test_scripts_fetch_content.py b/test/test_scripts_fetch_content.py index e3320da75..61cfb5b3a 100644 --- a/test/test_scripts_fetch_content.py +++ b/test/test_scripts_fetch_content.py @@ -1,3 +1,7 @@ +import bz2 +import gzip +import hashlib +import http.server import io import json import lzma @@ -7,7 +11,9 @@ import stat import sys import tarfile +import threading import urllib.request +import zipfile from importlib.machinery import SourceFileLoader from importlib.util import module_from_spec, spec_from_loader from unittest.mock import MagicMock @@ -35,6 +41,114 @@ def fetch_content_mod(): return mod +class RangeServer(http.server.ThreadingHTTPServer): + """Serves a single blob, optionally honouring range requests. + + ``faults`` maps a Range header value to a list of faults to inject, one + per request for that range: "error" answers 500, "truncate" sends half of + the body and closes the connection, and "slow" pauses half way through. + """ + + daemon_threads = True + allow_reuse_address = True + + def __init__( + self, content, ranges=True, accept_ranges="bytes", gzip=False, log=None + ): + super().__init__(("127.0.0.1", 0), RangeHandler) + self.content = content + self.ranges = ranges + self.accept_ranges = accept_ranges + self.gzip = gzip + self.faults = {} + self.requests = [] + self.log = log + self.lock = threading.Lock() + + @property + def url(self): + return "http://{}:{}/blob".format(*self.server_address) + + +class RangeHandler(http.server.BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def log_message(self, *args): + pass + + def do_GET(self): + content = self.server.content + requested = self.headers.get("Range") + with self.server.lock: + self.server.requests.append(requested) + if self.server.log is not None: + self.server.log.append((self.server, requested)) + faults = self.server.faults.get(requested) + fault = faults.pop(0) if faults else None + + if fault == "error": + self.send_response(500) + self.send_header("Content-Length", "0") + self.end_headers() + return + + start, end = 0, len(content) - 1 + partial = False + if requested and self.server.ranges: + start, _, last = requested.partition("=")[2].partition("-") + start, end = int(start), int(last) + if start >= len(content): + self.send_response(416) + self.send_header("Content-Range", f"bytes */{len(content)}") + self.send_header("Content-Length", "0") + self.end_headers() + return + end = min(end, len(content) - 1) + partial = True + + body = content[start : end + 1] + self.send_response(206 if partial else 200) + if partial: + self.send_header("Content-Range", f"bytes {start}-{end}/{len(content)}") + elif self.server.gzip and "gzip" in self.headers.get("Accept-Encoding", ""): + body = gzip.compress(body) + self.send_header("Content-Encoding", "gzip") + if self.server.accept_ranges: + self.send_header("Accept-Ranges", self.server.accept_ranges) + self.send_header("Content-Length", str(len(body))) + self.end_headers() + if fault == "truncate": + self.wfile.write(body[: len(body) // 2]) + self.close_connection = True + return + if fault == "slow": + self.wfile.write(body[: len(body) // 2]) + self.wfile.flush() + threading.Event().wait(1) + self.wfile.write(body[len(body) // 2 :]) + return + self.wfile.write(body) + + +@pytest.fixture +def serve(): + servers = [] + + def inner(content, **kwargs): + server = RangeServer(content, **kwargs) + threading.Thread( + target=server.serve_forever, kwargs={"poll_interval": 0.01}, daemon=True + ).start() + servers.append(server) + return server + + yield inner + + for server in servers: + server.shutdown() + server.server_close() + + @pytest.mark.parametrize( "url,sha256,size,headers,raises", ( @@ -276,7 +390,30 @@ def _make_tar(path, files): tar.addfile(info, io.BytesIO(data)) -def test_fetch_urls_merges_staged_extractions(monkeypatch, tmp_path, fetch_content_mod): +@pytest.fixture(params=("stream", "file")) +def mock_downloads(request, monkeypatch, fetch_content_mod): + """Serve downloads from a dict of names to content, streamed or not.""" + monkeypatch.setenv( + "TASKGRAPH_FETCH_STREAM", "1" if request.param == "stream" else "0" + ) + contents = {} + + def mock_download_to_path(url, path, sha256=None, size=None, headers=None): + path.write_bytes(contents[url.rsplit("/", 1)[-1]]) + + def mock_stream_download(url, sha256=None, size=None, headers=None): + data = contents[url.rsplit("/", 1)[-1]] + for i in range(0, len(data), 1000): + yield data[i : i + 1000] + + monkeypatch.setattr(fetch_content_mod, "download_to_path", mock_download_to_path) + monkeypatch.setattr(fetch_content_mod, "stream_download", mock_stream_download) + return contents + + +def test_fetch_urls_merges_staged_extractions( + tmp_path, fetch_content_mod, mock_downloads +): archives = tmp_path / "archives" archives.mkdir() _make_tar( @@ -289,11 +426,8 @@ def test_fetch_urls_merges_staged_extractions(monkeypatch, tmp_path, fetch_conte ) dest = tmp_path / "fetches" dest.mkdir() - - def mock_download_to_path(url, path, sha256=None, size=None): - shutil.copy(archives / path.name, path) - - monkeypatch.setattr(fetch_content_mod, "download_to_path", mock_download_to_path) + for name in ("common.tar", "suite.tar"): + mock_downloads[name] = (archives / name).read_bytes() fetch_content_mod.fetch_urls( [ @@ -313,20 +447,16 @@ def mock_download_to_path(url, path, sha256=None, size=None): assert (dest / "tests" / "shared.txt").read_text() == "second" -def test_fetch_urls_places_unextracted_files(monkeypatch, tmp_path, fetch_content_mod): +def test_fetch_urls_places_unextracted_files( + tmp_path, fetch_content_mod, mock_downloads +): archives = tmp_path / "archives" archives.mkdir() _make_tar(archives / "tool.tar", {"tool/bin/tool": "t"}) dest = tmp_path / "fetches" dest.mkdir() - - def mock_download_to_path(url, path, sha256=None, size=None): - if path.name.endswith(".tar"): - shutil.copy(archives / path.name, path) - else: - path.write_text("plain") - - monkeypatch.setattr(fetch_content_mod, "download_to_path", mock_download_to_path) + mock_downloads["tool.tar"] = (archives / "tool.tar").read_bytes() + mock_downloads["plain.txt"] = mock_downloads["notatar.txt"] = b"plain" fetch_content_mod.fetch_urls( [ @@ -472,3 +602,362 @@ def test_extract_archive_zstd_tar_failure(tmp_path, fetch_content_mod): with pytest.raises(Exception, match="exited"): fetch_content_mod.extract_archive(archive, dest) + + +def _compress(data, compression): + if compression == "zst": + zstandard = pytest.importorskip("zstandard") + return zstandard.ZstdCompressor(write_checksum=True).compress(data) + return { + "tar": lambda d: d, + "gz": gzip.compress, + "xz": lzma.compress, + "bz2": bz2.compress, + }[compression](data) + + +@pytest.fixture +def streamed(monkeypatch, fetch_content_mod): + """Enable streaming, and make sure nothing goes through a file.""" + monkeypatch.setenv("TASKGRAPH_FETCH_STREAM", "1") + monkeypatch.setattr(fetch_content_mod.time, "sleep", lambda s: None) + downloaded = [] + + def recording_download_to_path(url, path, *args, **kwargs): + downloaded.append(path.name) + return real_download_to_path(url, path, *args, **kwargs) + + real_download_to_path = fetch_content_mod.download_to_path + monkeypatch.setattr( + fetch_content_mod, "download_to_path", recording_download_to_path + ) + return downloaded + + +def _archive_url(server, name): + return server.url.replace("/blob", f"/{name}") + + +STREAM_FILES = { + "dir/a.txt": "a", + "dir/sub/b.txt": "b" * 5000, + # Big and incompressible enough to span many chunks. + "big.txt": os.urandom(300000).hex(), +} + + +def _assert_extracted(dest, files=STREAM_FILES): + extracted = { + p.relative_to(dest).as_posix(): p.read_text() + for p in dest.rglob("*") + if p.is_file() + } + assert extracted == files + + +@pytest.mark.skipif(sys.platform == "win32", reason="Windows extracts with tarfile") +@pytest.mark.parametrize( + "compression,commands", + ( + pytest.param("zst", [["zstd", "-dcq"], ["tar", "xf", "-"]]), + pytest.param("gz", [["tar", "xf", "-"]]), + pytest.param("xz", [["tar", "xf", "-"]]), + pytest.param("bz2", [["tar", "xf", "-"]]), + pytest.param("tar", [["tar", "xf", "-"]]), + ), +) +def test_stream_extract( + tmp_path, fetch_content_mod, serve, streamed, popen_calls, compression, commands +): + if compression == "zst" and not shutil.which("zstd"): + pytest.skip("needs the zstd program") + data = _compress(_tar_bytes(tmp_path, STREAM_FILES), compression) + server = serve(data) + dest = tmp_path / "fetches" + dest.mkdir() + + fetch_content_mod.fetch_urls( + [ + ( + _archive_url(server, f"archive.tar.{compression}"), + dest, + True, + hashlib.sha256(data).hexdigest(), + len(data), + ) + ] + ) + + _assert_extracted(dest) + assert server.requests == [None] + assert streamed == [] + assert popen_calls == commands + + +@pytest.mark.skipif(sys.platform == "win32", reason="Windows extracts with tarfile") +def test_stream_extract_zstd_without_zstd_program( + tmp_path, fetch_content_mod, serve, streamed, popen_calls, monkeypatch +): + data = _compress(_tar_bytes(tmp_path, STREAM_FILES), "zst") + server = serve(data) + monkeypatch.setattr(fetch_content_mod.shutil, "which", lambda name: None) + dest = tmp_path / "dest" + dest.mkdir() + + fetch_content_mod.fetch_and_extract(_archive_url(server, "a.tar.zst"), dest) + + _assert_extracted(dest) + assert streamed == [] + assert popen_calls == [["tar", "xf", "-"]] + + +def test_stream_extract_windows( + tmp_path, fetch_content_mod, serve, streamed, popen_calls, monkeypatch +): + data = _compress(_tar_bytes(tmp_path, STREAM_FILES), "gz") + server = serve(data) + monkeypatch.setattr(fetch_content_mod.sys, "platform", "win32") + dest = tmp_path / "dest" + dest.mkdir() + + fetch_content_mod.fetch_and_extract(_archive_url(server, "a.tar.gz"), dest) + + _assert_extracted(dest) + assert streamed == [] + assert popen_calls == [] + + +def test_stream_extract_retry(tmp_path, fetch_content_mod, serve, streamed): + """A stream that breaks off is extracted again from scratch.""" + data = _compress(_tar_bytes(tmp_path, STREAM_FILES), "gz") + server = serve(data) + server.faults[None] = ["truncate"] + dest = tmp_path / "fetches" + dest.mkdir() + + fetch_content_mod.fetch_urls( + [(_archive_url(server, "a.tar.gz"), dest, True, None, None)] + ) + + _assert_extracted(dest) + assert server.requests == [None, None] + assert streamed == [] + + +def test_stream_extract_partial_output_not_merged( + tmp_path, fetch_content_mod, serve, streamed, monkeypatch +): + data = _compress(_tar_bytes(tmp_path, STREAM_FILES), "tar") + server = serve(data) + server.faults[None] = ["truncate"] * 5 + dest = tmp_path / "fetches" + dest.mkdir() + emptied = [] + real_empty_directory = fetch_content_mod.empty_directory + + def recording_empty_directory(path): + # Half of the archive made it, so some files were extracted. + emptied.append(sorted(p.name for p in pathlib.Path(path).rglob("*"))) + real_empty_directory(path) + + monkeypatch.setattr(fetch_content_mod, "empty_directory", recording_empty_directory) + + with pytest.raises(Exception, match="no more retries"): + fetch_content_mod.fetch_urls( + [(_archive_url(server, "a.tar"), dest, True, None, None)] + ) + + assert len(emptied) == 5 + assert all("a.txt" in names for names in emptied), emptied + # Only the empty staging directory is left behind. + (staging,) = dest.iterdir() + assert staging.name.startswith(".fetch.") + assert list(staging.iterdir()) == [] + + +@pytest.mark.parametrize("mismatch", ("sha256", "size")) +def test_stream_extract_integrity_mismatch( + tmp_path, fetch_content_mod, serve, streamed, mismatch +): + data = _compress(_tar_bytes(tmp_path, STREAM_FILES), "gz") + server = serve(data) + dest = tmp_path / "fetches" + dest.mkdir() + sha256 = "0" * 64 if mismatch == "sha256" else None + size = len(data) + 1 if mismatch == "size" else None + + with pytest.raises(Exception, match="no more retries"): + fetch_content_mod.fetch_urls( + [(_archive_url(server, "a.tar.gz"), dest, True, sha256, size)] + ) + + assert len(server.requests) == 5 + assert [p.name for p in dest.rglob("*") if not p.name.startswith(".fetch.")] == [] + + +@pytest.mark.parametrize("compression", ("gz", "zst")) +def test_stream_extract_broken_archive_not_retried( + tmp_path, fetch_content_mod, serve, streamed, compression +): + """An intact download of a broken archive fails without retrying.""" + good = _compress(_tar_bytes(tmp_path, STREAM_FILES), compression) + # Corrupt the middle, well after the start used to detect the type. + mid = len(good) // 2 + data = ( + good[:mid] + + bytes(b ^ 0xFF for b in good[mid : mid + 2000]) + + good[mid + 2000 :] + ) + server = serve(data) + dest = tmp_path / "dest" + dest.mkdir() + + with pytest.raises(Exception) as excinfo: + fetch_content_mod.fetch_and_extract( + _archive_url(server, f"a.tar.{compression}"), + dest, + sha256=hashlib.sha256(data).hexdigest(), + ) + + assert "no more retries" not in str(excinfo.value) + assert server.requests == [None] + assert list(dest.iterdir()) == [] + + +def test_stream_extract_not_an_archive(tmp_path, fetch_content_mod, serve, streamed): + """Something that isn't a tar is written out whole, from one request.""" + data = os.urandom(3 * 1024 * 1024) + server = serve(data) + dest = tmp_path / "dest" + dest.mkdir() + + fetch_content_mod.fetch_and_extract( + _archive_url(server, "tool.exe"), dest, sha256=hashlib.sha256(data).hexdigest() + ) + + assert (dest / "tool.exe").read_bytes() == data + assert server.requests == [None] + assert streamed == [] + + +@pytest.mark.skipif(not shutil.which("unzip"), reason="needs unzip") +def test_zip_goes_through_file(tmp_path, fetch_content_mod, serve, streamed): + archive = tmp_path / "archive.zip" + with zipfile.ZipFile(archive, "w") as zf: + zf.writestr("dir/a.txt", "a") + server = serve(archive.read_bytes()) + dest = tmp_path / "dest" + dest.mkdir() + + fetch_content_mod.fetch_and_extract(_archive_url(server, "archive.zip"), dest) + + assert streamed == ["archive.zip"] + assert (dest / "dir" / "a.txt").read_text() == "a" + assert not (dest / "archive.zip").exists() + + +def test_no_extract_goes_through_file(tmp_path, fetch_content_mod, serve, streamed): + data = _compress(_tar_bytes(tmp_path, STREAM_FILES), "gz") + server = serve(data) + dest = tmp_path / "dest" + dest.mkdir() + + fetch_content_mod.fetch_and_extract( + _archive_url(server, "a.tar.gz"), dest, extract=False + ) + + assert streamed == ["a.tar.gz"] + assert (dest / "a.tar.gz").read_bytes() == data + + +def test_stream_disabled(tmp_path, fetch_content_mod, serve, streamed, monkeypatch): + data = _compress(_tar_bytes(tmp_path, STREAM_FILES), "gz") + server = serve(data) + monkeypatch.setenv("TASKGRAPH_FETCH_STREAM", "0") + dest = tmp_path / "dest" + dest.mkdir() + + fetch_content_mod.fetch_and_extract(_archive_url(server, "a.tar.gz"), dest) + + assert streamed == ["a.tar.gz"] + _assert_extracted(dest) + + +@pytest.mark.parametrize( + "value,expected", + ( + pytest.param(None, True, id="unset"), + pytest.param("1", True, id="1"), + pytest.param("0", False, id="0"), + pytest.param("false", False, id="false"), + ), +) +def test_streaming_enabled(fetch_content_mod, monkeypatch, value, expected): + monkeypatch.delenv("TASKGRAPH_FETCH_STREAM", raising=False) + if value is not None: + monkeypatch.setenv("TASKGRAPH_FETCH_STREAM", value) + assert fetch_content_mod.streaming_enabled() is expected + + +def test_iter_reader_iter_from_start(fetch_content_mod): + chunks = [b"abc", b"defg", b"hi"] + reader = fetch_content_mod.IterReader(iter(chunks)) + fh = io.BufferedReader(reader, 2) + assert fh.read(5) == b"abcde" + assert b"".join(reader.iter_from_start()) == b"abcdefghi" + + +def test_iter_reader_failure(fetch_content_mod): + def chunks(): + yield b"abc" + raise OSError("connection reset") + + reader = fetch_content_mod.IterReader(chunks()) + assert reader.read(10) == b"abc" + assert not reader.failed + with pytest.raises(OSError): + reader.drain() + assert reader.failed + + +@pytest.mark.skipif( + sys.platform == "win32" or os.getuid() == 0, reason="needs POSIX directory modes" +) +def test_empty_directory_readonly(tmp_path, fetch_content_mod): + (tmp_path / "ro" / "sub").mkdir(parents=True) + (tmp_path / "ro" / "sub" / "f").write_text("f") + outside = tmp_path.parent / f"{tmp_path.name}-outside" + outside.mkdir() + (tmp_path / "link").symlink_to(outside) + (tmp_path / "ro" / "sub").chmod(0o555) + (tmp_path / "ro").chmod(0o555) + + fetch_content_mod.empty_directory(tmp_path) + + assert list(tmp_path.iterdir()) == [] + assert outside.exists() + + +def test_stream_extract_verifies_unread_tail( + tmp_path, fetch_content_mod, serve, streamed, monkeypatch +): + """tarfile stops reading at the end of the archive, which may come well + before the end of the download. The whole download must still get + verified.""" + # Trailing zeros are valid tar padding. + data = _tar_bytes(tmp_path, STREAM_FILES) + bytes(5 * 1024 * 1024) + server = serve(data) + monkeypatch.setattr(fetch_content_mod.sys, "platform", "win32") + dest = tmp_path / "dest" + dest.mkdir() + + with pytest.raises(Exception, match="no more retries"): + fetch_content_mod.fetch_and_extract( + _archive_url(server, "a.tar"), dest, sha256="0" * 64 + ) + assert list(dest.iterdir()) == [] + + fetch_content_mod.fetch_and_extract( + _archive_url(server, "a.tar"), dest, sha256=hashlib.sha256(data).hexdigest() + ) + _assert_extracted(dest)