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 337d807fb..6c8390f8c 100755 --- a/src/taskgraph/run-task/fetch-content +++ b/src/taskgraph/run-task/fetch-content @@ -8,9 +8,11 @@ import bz2 import concurrent.futures import contextlib import datetime +import functools import gzip import hashlib import io +import itertools import json import lzma import multiprocessing @@ -43,6 +45,12 @@ 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 + def log(msg): print(msg, file=sys.stderr) @@ -219,7 +227,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 @@ -248,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.""" @@ -336,12 +351,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": @@ -350,7 +372,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: @@ -365,7 +387,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) == ZSTD_MAGIC def archive_type(path: pathlib.Path): @@ -466,6 +493,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 @@ -478,23 +512,211 @@ 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 - p.stdin.write(chunk) +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) + + 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") + +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( @@ -719,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 a971c22c5..61cfb5b3a 100644 --- a/test/test_scripts_fetch_content.py +++ b/test/test_scripts_fetch_content.py @@ -1,12 +1,19 @@ +import bz2 +import gzip +import hashlib +import http.server import io import json +import lzma import os import pathlib import shutil 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 @@ -34,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", ( @@ -275,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( @@ -288,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( [ @@ -312,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( [ @@ -379,3 +510,454 @@ 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) + + +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)