From ed88002c3783465254117c7b5df2c3a9f59b4784 Mon Sep 17 00:00:00 2001 From: "J. Nick Koston" Date: Wed, 30 Sep 2026 02:23:05 +0200 Subject: [PATCH] [core] Extract native toolchain package archives in parallel (#18840) --- esphome/arduino8266/framework.py | 23 +- esphome/espidf/framework.py | 4 +- esphome/framework_helpers.py | 100 ++++-- esphome/platformio/library.py | 4 +- esphome/platformio/prefetch.py | 9 +- esphome/platformio/registry.py | 155 +++++++++- .../unit_tests/test_arduino8266_framework.py | 51 +--- tests/unit_tests/test_framework_helpers.py | 91 +++++- tests/unit_tests/test_platformio_prefetch.py | 2 +- tests/unit_tests/test_platformio_registry.py | 289 ++++++++++++++++-- 10 files changed, 591 insertions(+), 137 deletions(-) diff --git a/esphome/arduino8266/framework.py b/esphome/arduino8266/framework.py index a8f8c65c06..b731b74675 100644 --- a/esphome/arduino8266/framework.py +++ b/esphome/arduino8266/framework.py @@ -24,9 +24,10 @@ from esphome.core import EsphomeError, Version from esphome.framework_helpers import str_to_lst_of_str from esphome.platformio.registry import ( Download, + PackageSpec, Resolver, get_systype, - install_package, + install_packages, prefetch_packages, ) @@ -151,14 +152,14 @@ def check_and_install(framework_version: Version) -> InstalledPaths: toolchain_path = get_toolchain_path() # One spec per package: the prefetch and the installs must agree specs = ( - ( + PackageSpec( FRAMEWORK_PACKAGE, release.tag, framework_path, ESPHOME_ARDUINO8266_FRAMEWORK_MIRRORS, ("cores/esp8266", "tools/sdk", "libraries"), ), - ( + PackageSpec( TOOLCHAIN_PACKAGE, TOOLCHAIN_VERSION, toolchain_path, @@ -174,18 +175,10 @@ def check_and_install(framework_version: Version) -> InstalledPaths: resolvers[FRAMEWORK_PACKAGE] = release.download if not ESPHOME_ARDUINO8266_TOOLCHAIN_MIRRORS: resolvers[TOOLCHAIN_PACKAGE] = toolchain_download - # Fetch both archives at once; the installs below verify and extract - prefetch_packages([spec[:4] for spec in specs], downloads_dir, resolvers) - for name, version, dest, mirrors, expect in specs: - install_package( - name, - version, - dest, - mirrors, - downloads_dir, - expect=expect, - resolve=resolvers.get(name), - ) + # Fetch both archives at once; the installs verify and extract them. + # One spec list for both, so the two phases cannot drift. + prefetch_packages(specs, downloads_dir, resolvers) + install_packages(specs, downloads_dir, resolvers) return InstalledPaths( framework=framework_path, toolchain=toolchain_path, ninja=ninja_path ) diff --git a/esphome/espidf/framework.py b/esphome/espidf/framework.py index 1a4b17efdb..8c377561ca 100644 --- a/esphome/espidf/framework.py +++ b/esphome/espidf/framework.py @@ -34,7 +34,7 @@ from esphome.framework_helpers import ( run_command_ok, str_to_lst_of_str, tool_version_runs, - warn_prefetch_failures, + warn_batch_failures, ) from esphome.helpers import write_file_if_changed @@ -774,7 +774,7 @@ def _prefetch_idf_tool_archives( for entry in entries ], ) - warn_prefetch_failures(failures) + warn_batch_failures(failures, "Could not prefetch %s: %s") if len(failures) == len(entries): # A systematic fault, not one flaky mirror: the resume # workaround (#17703) is off for this whole install diff --git a/esphome/framework_helpers.py b/esphome/framework_helpers.py index fc2a18a6ec..86010e3065 100644 --- a/esphome/framework_helpers.py +++ b/esphome/framework_helpers.py @@ -15,7 +15,7 @@ import threading import time from typing import IO, TYPE_CHECKING -from esphome.helpers import ProgressBar, rmtree +from esphome.helpers import ProgressBar, get_usable_cpu_count, rmtree from esphome.net_retry import ( NETWORK_MAX_ATTEMPTS, http_request, @@ -288,10 +288,25 @@ def _detect_archive_root(names: Iterable[str]) -> str | None: return root if has_descendant else None +def _resolve_progress( + progress: Callable[[float], None] | None, + progress_header: str | None, + has_work: bool, +) -> Callable[[float], None] | None: + """Fraction reporter for an extractor: the caller's callback wins over a + private ``progress_header`` bar.""" + if progress is not None: + return progress + if progress_header and has_work: + return ProgressBar(progress_header).update + return None + + def _tar_extract_all( data: io.BufferedIOBase, extract_dir: PathType = ".", progress_header: str | None = None, + progress: Callable[[float], None] | None = None, ): """ Extract a TAR archive to the specified directory. @@ -306,6 +321,7 @@ def _tar_extract_all( data: File-like object containing the TAR archive extract_dir: Directory to extract contents to progress_header: If set, show a progress bar with this header + progress: fraction callback (0..1, ends at 1.0); overrides progress_header """ import tarfile @@ -364,21 +380,23 @@ def _tar_extract_all( safe_members.append(member) total = len(safe_members) - progress = ( - ProgressBar(progress_header) if progress_header and total > 0 else None - ) + report = _resolve_progress(progress, progress_header, total > 0) for i, member in enumerate(safe_members, 1): - tar_ref.extract(member, abs_dest) - if progress is not None: - progress.update(i / total) - if progress is not None: - progress.update(1) + # Named: the default is fully_trusted on 3.12/3.13, data on + # 3.14. The pre-pass drops unsafe members; an escape past it + # gains nothing, since the build runs what these archives hold. + tar_ref.extract(member, abs_dest, filter="fully_trusted") + if report is not None: + report(i / total) + if report is not None: + report(1) def _zip_extract_all( data: io.BufferedIOBase, extract_dir: PathType = ".", progress_header: str | None = None, + progress: Callable[[float], None] | None = None, ): """ Extract a ZIP archive to the specified directory. @@ -387,6 +405,7 @@ def _zip_extract_all( data: File-like object containing the ZIP archive extract_dir: Directory to extract contents to progress_header: If set, show a progress bar with this header + progress: fraction callback (0..1, ends at 1.0); overrides progress_header """ import zipfile @@ -403,9 +422,7 @@ def _zip_extract_all( strip_prefix = f"{strip_root}/" if strip_root is not None else None total = len(all_members) - progress = ( - ProgressBar(progress_header) if progress_header and total > 0 else None - ) + report = _resolve_progress(progress, progress_header, total > 0) for i, member in enumerate(all_members, 1): # 1. Normalize name @@ -438,10 +455,10 @@ def _zip_extract_all( # 6. Extract zip_ref.extract(member, extract_dir) - if progress is not None: - progress.update(i / total) - if progress is not None: - progress.update(1) + if report is not None: + report(i / total) + if report is not None: + report(1) def _rename_with_retry( @@ -472,6 +489,7 @@ def _7z_extract_all( data: io.BufferedIOBase, extract_dir: PathType = ".", progress_header: str | None = None, + progress: Callable[[float], None] | None = None, ): """ Extract a 7z archive to the specified directory. @@ -486,6 +504,7 @@ def _7z_extract_all( data: File-like object containing the 7z archive (must be seekable) extract_dir: Directory to extract contents to progress_header: If set, show a progress bar with this header + progress: called with 1.0 on completion; overrides progress_header """ import py7zr @@ -524,19 +543,15 @@ def _7z_extract_all( continue safe_targets.append(raw) - progress = ( - ProgressBar(progress_header) - if progress_header and safe_targets - else None - ) + report = _resolve_progress(progress, progress_header, bool(safe_targets)) if len(safe_targets) == len(all_names): z.extractall(path=staging) else: z.extract(path=staging, targets=safe_targets) - if progress is not None: - progress.update(1) + if report is not None: + report(1) src_root = staging / strip_root if strip_root else staging for item in src_root.iterdir(): @@ -567,6 +582,7 @@ def archive_extract_all( archive: PathType | io.RawIOBase | IO[bytes], extract_dir: PathType = ".", progress_header: str | None = None, + progress: Callable[[float], None] | None = None, ): """ Extract an archive file to the specified directory. @@ -575,6 +591,7 @@ def archive_extract_all( archive: Path to archive file or file-like object extract_dir: Directory to extract contents to progress_header: If set, show a progress bar with this header + progress: fraction callback (0..1, ends at 1.0); overrides progress_header Raises: TypeError: If archive is not a valid type @@ -605,7 +622,9 @@ def archive_extract_all( break if matched_fct is None: raise ValueError("Unsupported archive format") - matched_fct(archive_ref, extract_dir, progress_header=progress_header) + matched_fct( + archive_ref, extract_dir, progress_header=progress_header, progress=progress + ) def _open_ranged( @@ -769,13 +788,23 @@ def _stream_response_to_file( # hammering the host or the mirrors. BATCH_DOWNLOAD_WORKERS = 4 +# Measured: gz peaks near 2 workers (8 is slower than serial), xz +# plateaus by 4 and holds ~50 MB of dictionary per worker. +BATCH_EXTRACT_WORKERS = 4 + + +def extract_workers(jobs: int | None = None) -> int: + """Worker count for an extraction batch of ``jobs`` archives.""" + workers = min(get_usable_cpu_count(), BATCH_EXTRACT_WORKERS) + return workers if jobs is None else min(workers, jobs) + def run_batch_downloads( header: str, jobs: list[tuple[str, int, Callable[[Callable[[int], None]], None]]], max_workers: int = BATCH_DOWNLOAD_WORKERS, ) -> list[tuple[str, BaseException]]: - """Run ``(name, size, fetch)`` download jobs concurrently under one bar. + """Run ``(name, size, fetch)`` jobs concurrently under one bar. Each ``fetch(tracker)`` reports absolute byte counts; the bar total is the sum of the sizes. Failures are returned after the bar is done so @@ -1005,15 +1034,26 @@ def resume_fetch_job( return fetch -def warn_prefetch_failures( +def is_expected_fetch_error(err: BaseException) -> bool: + """Download failures the callers degrade on, vs programming errors.""" + from esphome.core import EsphomeError # local import avoids circular dependency + + return isinstance(err, (EsphomeError, OSError)) + + +def warn_batch_failures( failures: list[tuple[str, BaseException]], - message: str = "Could not prefetch %s: %s", + message: str, ) -> None: - """Warn per failed batch-prefetch job; the caller's installer retries them.""" + """Warn per failed batch job, keeping the traceback of unexpected errors.""" for name, err in failures: # failure_reason: a message-less exception must not log blank - _LOGGER.warning(message, name, failure_reason(err)) - _LOGGER.debug("Prefetch failure detail", exc_info=err) + if is_expected_fetch_error(err): + _LOGGER.warning(message, name, failure_reason(err)) + _LOGGER.debug("Failure detail", exc_info=err) + else: + # A programming error must not be reduced to a bare message + _LOGGER.warning(message, name, failure_reason(err), exc_info=err) def download_with_resume( diff --git a/esphome/platformio/library.py b/esphome/platformio/library.py index e551d8f1c0..e5e4aa7245 100644 --- a/esphome/platformio/library.py +++ b/esphome/platformio/library.py @@ -35,7 +35,7 @@ from esphome.framework_helpers import ( failure_reason, rmdir, run_batch_downloads, - warn_prefetch_failures, + warn_batch_failures, ) _LOGGER = logging.getLogger(__name__) @@ -1109,7 +1109,7 @@ def _prefetch_wave( + [(c.name, 0, partial(_clone_source, c, salt, namespace)) for c in clones], ) # The sequential call below retries and raises the real error - warn_prefetch_failures( + warn_batch_failures( failures, "Prefetch of %s failed (retrying sequentially): %s" ) except Exception as err: # noqa: BLE001 # pylint: disable=broad-exception-caught diff --git a/esphome/platformio/prefetch.py b/esphome/platformio/prefetch.py index e648192b73..fbb31ae452 100644 --- a/esphome/platformio/prefetch.py +++ b/esphome/platformio/prefetch.py @@ -37,13 +37,14 @@ from esphome.framework_helpers import ( content_length, discard_partial_download, downloaded_bytes, + extract_workers, failure_reason, resume_fetch_job, run_batch_downloads, wait_for_download_lock, - warn_prefetch_failures, + warn_batch_failures, ) -from esphome.helpers import get_bool_env, get_usable_cpu_count, rmtree +from esphome.helpers import get_bool_env, rmtree _LOGGER = logging.getLogger(__name__) @@ -702,7 +703,7 @@ def _preinstall( would hang, not fail). Waves skip dependencies; the installed manifests feed the next wave. Any failure falls back to pio run. """ - workers = min(get_usable_cpu_count(), len(entries)) + workers = extract_workers(len(entries)) # One manager per worker (_install mutates instance state); built # serially because construction rewires the shared manager logger managers: SimpleQueue = SimpleQueue() @@ -891,7 +892,7 @@ def _prefetch(build_dir: Path, env: str) -> None: ) # PlatformIO retries failed packages itself, without resume failures = run_batch_downloads("Downloading PlatformIO packages", jobs) - warn_prefetch_failures(failures) + warn_batch_failures(failures, "Could not prefetch %s: %s") failed_names = {name for name, _ in failures} elif not groups and not unresolved: # Record the no-work run so the parent skips the next spawn. diff --git a/esphome/platformio/registry.py b/esphome/platformio/registry.py index cab536c8da..326a587fc8 100644 --- a/esphome/platformio/registry.py +++ b/esphome/platformio/registry.py @@ -18,9 +18,12 @@ from esphome.framework_helpers import ( download_from_mirrors, download_with_resume, downloaded_bytes, + extract_workers, + is_expected_fetch_error, rmdir, run_batch_downloads, wait_for_download_lock, + warn_batch_failures, ) from esphome.net_retry import fetch_with_retry, http_request @@ -174,6 +177,16 @@ def _check_layout(name: str, dest: Path, expect: Collection[str]) -> None: ) +class PackageSpec(NamedTuple): + """One registry package to install.""" + + name: str + version: str + dest: Path + mirrors: list[str] + expect: Collection[str] = () + + class _PendingArchive(NamedTuple): name: str version: str @@ -194,27 +207,47 @@ def is_installed(dest: Path) -> bool: return (dest / ".esphome_extracted").is_file() +def _batched_download_progress( + name: str, version: str, extract_progress: Callable[[float], None] +) -> Callable[[int], None]: + """Zero-tick tracker for a batched install; announces a real download + once, since the shared bar cannot move for it.""" + ticks = 0 + + def progress(done: int) -> None: + nonlocal ticks + ticks += 1 + # A verified archive credits itself in one tick; more than one + # means bytes are streaming, including a resumed .part + if ticks == 2: + _LOGGER.info("Re-downloading %s %s ...", name, version) + extract_progress(0.0) + + return progress + + def prefetch_packages( - packages: list[tuple[str, str, Path, list[str]]], + packages: Collection[PackageSpec], downloads_dir: Path, resolvers: dict[str, Resolver] | None = None, ) -> None: """Download pending package archives in parallel under one combined bar. - ``packages`` holds ``(name, version, dest, mirrors)`` per package; - ``resolvers`` replaces the registry lookup by name. Purely - an optimization: ``install_package`` verifies every archive and - re-downloads anything this pass left unfinished. Mirror overrides and - registry entries without a size stay on the sequential path so its - per-file bars remain trustworthy. Each fetch holds the same per-dest - lock as ``install_package``: the archive's ``.part`` file is shared, and - two concurrent writers would truncate each other's bytes. + ``packages`` holds one ``PackageSpec`` per package, the same list the + install pass takes; ``expect`` is unused here and ``resolvers`` replaces + the registry lookup by name. Purely an optimization: ``install_package`` + verifies every archive and re-downloads anything this pass left + unfinished. Mirror overrides and registry entries without a size stay on + the sequential path so its per-file bars remain trustworthy. Each fetch + holds the same per-dest lock as ``install_package``: the archive's + ``.part`` file is shared, and two concurrent writers would truncate each + other's bytes. """ from filelock import FileLock, Timeout pending: list[_PendingArchive] = [] seen: set[Path] = set() - for name, version, dest, mirrors in packages: + for name, version, dest, mirrors, _expect in packages: if mirrors or is_installed(dest): continue archive = _archive_path(downloads_dir, name, version) @@ -283,7 +316,7 @@ def prefetch_packages( [(entry.name, entry.size, partial(_fetch, entry)) for entry in pending], ) for name, err in failures: - if isinstance(err, (EsphomeError, OSError)): + if is_expected_fetch_error(err): # Expected download failures: install_package retries this one # itself, with a visible bar _LOGGER.debug("Prefetch of %s failed: %s", name, err) @@ -301,6 +334,7 @@ def install_package( downloads_dir: Path, expect: Collection[str], resolve: Resolver | None = None, + extract_progress: Callable[[float], None] | None = None, ) -> None: """Download, verify, and extract one package if not already installed. @@ -309,6 +343,9 @@ def install_package( substitution) is trusted as configured. ``downloads_dir`` holds the archive between runs so an interrupted download resumes. ``resolve`` replaces the registry lookup. + + ``extract_progress`` receives extraction fractions in [0, 1] instead of + the private per-file bars (see ``install_packages``). """ if not expect: # Layout validation before marker.touch() is the only guard against @@ -332,7 +369,10 @@ def install_package( # Persistent location so an interrupted download resumes across runs. downloads_dir.mkdir(parents=True, exist_ok=True) archive = _archive_path(downloads_dir, name, version) - _LOGGER.info("Downloading %s %s ...", name, version) + # Batched runs are announced by the batch header + batched = extract_progress is not None and archive.is_file() + if not batched: + _LOGGER.info("Downloading %s %s ...", name, version) if mirrors: _LOGGER.warning( "Downloading %s from a mirror override; checksum verification " @@ -346,11 +386,96 @@ def install_package( url, sha256, size = ( resolve() if resolve else registry_download(name, version) ) - download_with_resume(url, archive, sha256=sha256, size=size) - _LOGGER.info("Extracting %s ...", name) - archive_extract_all(archive, dest, progress_header="Extracting") + download_with_resume( + url, + archive, + sha256=sha256, + size=size, + # Zero ticks: the shared bar must never run backwards + progress=None + if extract_progress is None + else _batched_download_progress(name, version, extract_progress), + ) + if not batched: + _LOGGER.info("Extracting %s ...", name) + archive_extract_all( + archive, dest, progress_header="Extracting", progress=extract_progress + ) # Validate the layout before recording success, so an unexpected # package is never cached as a working install. _check_layout(name, dest, expect) marker.touch() archive.unlink(missing_ok=True) + + +def install_packages( + specs: Collection[PackageSpec], + downloads_dir: Path, + resolvers: dict[str, Resolver] | None = None, +) -> None: + """Install several packages; prefetched archives extract in parallel under + one shared bar, the rest take the sequential ``install_package`` path. + ``resolvers`` replaces the registry lookup by name; the first failure is + re-raised.""" + resolvers = resolvers or {} + pending: list[tuple[PackageSpec, int]] = [] + rest: list[PackageSpec] = [] + for spec in specs: + name, version, dest, mirrors, _expect = spec + archive = _archive_path(downloads_dir, name, version) + if is_installed(dest) or mirrors: + rest.append(spec) + continue + try: + # Sized, not hashed: install_package still verifies the archive + size = archive.stat().st_size + except FileNotFoundError: + rest.append(spec) + continue + pending.append((spec, size)) + if len(pending) < 2: + # One archive alone gains nothing from a pool + rest = list(specs) + pending = [] + + def _install(spec: PackageSpec, size: int, tracker: Callable[[int], None]) -> None: + name, version, dest, mirrors, expect = spec + install_package( + name, + version, + dest, + mirrors, + downloads_dir, + expect=expect, + resolve=resolvers.get(name), + extract_progress=lambda frac: tracker(int(frac * size)), + ) + + if pending: + workers = extract_workers(len(pending)) + _LOGGER.info( + "Extracting %d package archive(s) with %d worker(s): %s", + len(pending), + workers, + ", ".join(spec.name for spec, _ in pending), + ) + failures = run_batch_downloads( + "Extracting packages", + [(spec[0], size, partial(_install, spec, size)) for spec, size in pending], + max_workers=workers, + ) + if failures: + # The raised exception may not name the package; nothing runs + # behind this pass to redo the work + warn_batch_failures(failures, "Could not install %s: %s") + raise failures[0][1] + for name, version, dest, mirrors, expect in rest: + install_package( + name, + version, + dest, + mirrors, + downloads_dir, + expect=expect, + resolve=resolvers.get(name), + ) diff --git a/tests/unit_tests/test_arduino8266_framework.py b/tests/unit_tests/test_arduino8266_framework.py index a4193b9a74..5ffcea0114 100644 --- a/tests/unit_tests/test_arduino8266_framework.py +++ b/tests/unit_tests/test_arduino8266_framework.py @@ -92,16 +92,13 @@ def test_check_and_install_mirror_skips_pinned_toolchain(tmp_path: Path) -> None patch.dict(os.environ, {"ESPHOME_ARDUINO8266_PREFIX": str(tmp_path)}), patch.object(framework, "ESPHOME_ARDUINO8266_FRAMEWORK_MIRRORS", ["http://f"]), patch.object(framework, "ESPHOME_ARDUINO8266_TOOLCHAIN_MIRRORS", ["http://m"]), - patch.object(framework, "install_package") as mock_install, + patch.object(framework, "install_packages") as mock_install, patch.object(framework, "prefetch_packages") as mock_prefetch, patch.object(framework, "find_ninja", return_value=tmp_path / "ninja"), ): framework.check_and_install(cv.Version(3, 1, 2)) assert mock_prefetch.call_args.args[2] == {} - assert [call.kwargs["resolve"] for call in mock_install.call_args_list] == [ - None, - None, - ] + assert mock_install.call_args.args[2] == {} def test_check_and_install_installed_toolchain_on_unsupported_host( @@ -142,7 +139,7 @@ def test_check_and_install_unsupported_host_without_toolchain_raises( def test_check_and_install_returns_paths(tmp_path: Path) -> None: with ( patch.dict(os.environ, {"ESPHOME_ARDUINO8266_PREFIX": str(tmp_path)}), - patch.object(framework, "install_package") as mock_install, + patch.object(framework, "install_packages") as mock_install, patch.object(framework, "prefetch_packages") as mock_prefetch, patch.object(framework, "find_ninja", return_value=tmp_path / "ninja"), ): @@ -150,54 +147,38 @@ def test_check_and_install_returns_paths(tmp_path: Path) -> None: assert paths.framework == tmp_path / "frameworks" / _recommended().tag assert paths.toolchain == tmp_path / "toolchains" / framework.TOOLCHAIN_VERSION assert paths.ninja == tmp_path / "ninja" - assert mock_install.call_count == 2 # Full argument pinning: a copy-paste swap between the two near-identical - # calls (mirrors, destination) must not stay green - fw_call, tc_call = mock_install.call_args_list - assert fw_call.args == ( - framework.FRAMEWORK_PACKAGE, - _recommended().tag, - tmp_path / "frameworks" / _recommended().tag, - framework.ESPHOME_ARDUINO8266_FRAMEWORK_MIRRORS, - tmp_path / "downloads", - ) - assert fw_call.kwargs == { - "expect": ("cores/esp8266", "tools/sdk", "libraries"), - "resolve": _recommended().download, - } - assert tc_call.args == ( - framework.TOOLCHAIN_PACKAGE, - framework.TOOLCHAIN_VERSION, - tmp_path / "toolchains" / framework.TOOLCHAIN_VERSION, - framework.ESPHOME_ARDUINO8266_TOOLCHAIN_MIRRORS, - tmp_path / "downloads", - ) - assert tc_call.kwargs == { - "expect": ("bin", "xtensa-lx106-elf"), - "resolve": framework.toolchain_download, - } - # The prefetch sees the same package specs as the installs - assert mock_prefetch.call_args.args == ( - [ + # specs (mirrors, destination) must not stay green + assert mock_install.call_args.args == ( + ( ( framework.FRAMEWORK_PACKAGE, _recommended().tag, tmp_path / "frameworks" / _recommended().tag, framework.ESPHOME_ARDUINO8266_FRAMEWORK_MIRRORS, + ("cores/esp8266", "tools/sdk", "libraries"), ), ( framework.TOOLCHAIN_PACKAGE, framework.TOOLCHAIN_VERSION, tmp_path / "toolchains" / framework.TOOLCHAIN_VERSION, framework.ESPHOME_ARDUINO8266_TOOLCHAIN_MIRRORS, + ("bin", "xtensa-lx106-elf"), ), - ], + ), tmp_path / "downloads", { framework.FRAMEWORK_PACKAGE: _recommended().download, framework.TOOLCHAIN_PACKAGE: framework.toolchain_download, }, ) + # One spec list feeds both phases, so they cannot drift + assert mock_prefetch.call_args.args == mock_install.call_args.args + # PackageSpec instances, not bare tuples: the batch header reads .name + assert all( + isinstance(spec, framework.PackageSpec) + for spec in mock_install.call_args.args[0] + ) def test_get_build_env_prepends_toolchain_bin(tmp_path: Path) -> None: diff --git a/tests/unit_tests/test_framework_helpers.py b/tests/unit_tests/test_framework_helpers.py index 22b34c9df5..f3b182073f 100644 --- a/tests/unit_tests/test_framework_helpers.py +++ b/tests/unit_tests/test_framework_helpers.py @@ -523,6 +523,17 @@ class TestArchiveExtractAll: archive_extract_all(archive, dest) assert (dest / "file.txt").read_text() == "hi" + def test_progress_callback_passed_through(self, tmp_path: Path) -> None: + """The progress kwarg reaches the dispatched extractor.""" + archive = tmp_path / "test.tar.gz" + archive.write_bytes(_gzip_tar_bytes({"file.txt": b"hello"})) + dest = tmp_path / "out" + dest.mkdir() + fractions: list[float] = [] + archive_extract_all(archive, dest, progress=fractions.append) + assert fractions[-1] == 1 + assert (dest / "file.txt").read_bytes() == b"hello" + def test_invalid_type_raises_type_error(self) -> None: with pytest.raises(TypeError, match="archive must be"): archive_extract_all(42, ".") # type: ignore[arg-type] @@ -1951,6 +1962,19 @@ class TestTarExtractAllBranches: mock_pb.assert_called_once_with("Extracting") mock_pb.return_value.update.assert_called() + def test_progress_callback_replaces_bar(self, tmp_path: Path) -> None: + """A progress callback wins over progress_header and ends at 1.0.""" + buf = _make_tar([_reg("a.txt"), _reg("b.txt")], {"a.txt": b"x", "b.txt": b"y"}) + fractions: list[float] = [] + with patch("esphome.framework_helpers.ProgressBar") as mock_pb: + _tar_extract_all( + buf, tmp_path, progress_header="Extracting", progress=fractions.append + ) + mock_pb.assert_not_called() + assert fractions == sorted(fractions) + assert fractions[-1] == 1 + assert (tmp_path / "a.txt").is_file() + # --------------------------------------------------------------------------- # _zip_extract_all — additional branch coverage @@ -1980,6 +2004,19 @@ class TestZipExtractAllBranches: mock_pb.assert_called_once_with("Unzipping") mock_pb.return_value.update.assert_called() + def test_progress_callback_replaces_bar(self, tmp_path: Path) -> None: + """A progress callback wins over progress_header and ends at 1.0.""" + buf = _make_zip([("a.txt", "aaa"), ("b.txt", "bbb")]) + fractions: list[float] = [] + with patch("esphome.framework_helpers.ProgressBar") as mock_pb: + _zip_extract_all( + buf, tmp_path, progress_header="Unzipping", progress=fractions.append + ) + mock_pb.assert_not_called() + assert fractions == sorted(fractions) + assert fractions[-1] == 1 + assert (tmp_path / "a.txt").is_file() + # --------------------------------------------------------------------------- # _rename_with_retry @@ -2137,6 +2174,20 @@ class TestSevenZipExtractAll: mock_pb.assert_called_once_with("Unpacking 7z") mock_pb.return_value.update.assert_called() + def test_progress_callback_replaces_bar(self, tmp_path: Path) -> None: + """A progress callback wins over progress_header; 7z reports 1.0 once.""" + buf = self._make_7z({"file.txt": b"x"}) + out = tmp_path / "out" + out.mkdir() + fractions: list[float] = [] + with patch("esphome.framework_helpers.ProgressBar") as mock_pb: + _7z_extract_all( + buf, out, progress_header="Unpacking 7z", progress=fractions.append + ) + mock_pb.assert_not_called() + assert fractions == [1] + assert (out / "file.txt").is_file() + def test_absolute_path_in_names_skipped(self, tmp_path: Path) -> None: """Names that resolve as absolute are silently skipped.""" import py7zr @@ -2294,18 +2345,50 @@ def test_resume_fetch_job_threads_tracker(tmp_path: Path) -> None: ) -def test_warn_prefetch_failures_names_each_failure( +def test_warn_batch_failures_names_each_failure( caplog: pytest.LogCaptureFixture, ) -> None: """The shared failure loop warns per job with the failure reason.""" - from esphome.framework_helpers import warn_prefetch_failures + from esphome.framework_helpers import warn_batch_failures - warn_prefetch_failures([("toolchain-x@1", OSError("down"))]) + warn_batch_failures( + [("toolchain-x@1", OSError("down"))], "Could not prefetch %s: %s" + ) assert "Could not prefetch toolchain-x@1: down" in caplog.text - warn_prefetch_failures([("lib", OSError("gone"))], "Prefetch of %s failed: %s") + warn_batch_failures([("lib", OSError("gone"))], "Prefetch of %s failed: %s") assert "Prefetch of lib failed: gone" in caplog.text +def test_extract_workers_caps_and_clamps() -> None: + """Extraction stops scaling well before high core counts, and a batch + never asks for more workers than it has archives.""" + from esphome.framework_helpers import BATCH_EXTRACT_WORKERS, extract_workers + + with patch("esphome.framework_helpers.get_usable_cpu_count", return_value=64): + assert extract_workers() == BATCH_EXTRACT_WORKERS + assert extract_workers(2) == 2 + with patch("esphome.framework_helpers.get_usable_cpu_count", return_value=1): + assert extract_workers(8) == 1 + + +def test_warn_batch_failures_unexpected_error_keeps_traceback( + caplog: pytest.LogCaptureFixture, +) -> None: + """An unexpected error type is not reduced to a bare message; expected + download failures stay message-only at WARNING.""" + from esphome.framework_helpers import warn_batch_failures + + with caplog.at_level(logging.DEBUG): + warn_batch_failures( + [("pkg", TypeError("bad call")), ("lib", OSError("down"))], + "Could not install %s: %s", + ) + warnings = {r.getMessage(): r for r in caplog.records if r.levelname == "WARNING"} + assert warnings["Could not install pkg: bad call"].exc_info is not None + assert warnings["Could not install lib: down"].exc_info is None + assert "Failure detail" in caplog.text + + @pytest.mark.parametrize( ("platform", "input_path", "expected"), [ diff --git a/tests/unit_tests/test_platformio_prefetch.py b/tests/unit_tests/test_platformio_prefetch.py index 774493ecf4..ef573767f7 100644 --- a/tests/unit_tests/test_platformio_prefetch.py +++ b/tests/unit_tests/test_platformio_prefetch.py @@ -1727,7 +1727,7 @@ def test_preinstall_uses_distinct_managers_in_parallel(tmp_path: Path) -> None: barrier.wait() seed = _WaveManager(str(tmp_path)) - with patch.object(pf, "get_usable_cpu_count", return_value=2): + with patch("esphome.framework_helpers.get_usable_cpu_count", return_value=2): pf._preinstall( seed, [ diff --git a/tests/unit_tests/test_platformio_registry.py b/tests/unit_tests/test_platformio_registry.py index 9f1c140846..c30cdc7d6c 100644 --- a/tests/unit_tests/test_platformio_registry.py +++ b/tests/unit_tests/test_platformio_registry.py @@ -2,8 +2,10 @@ from __future__ import annotations -from contextlib import contextmanager +from collections.abc import Callable, Iterator +from contextlib import AbstractContextManager, contextmanager import json +import logging import os from pathlib import Path from unittest.mock import MagicMock, patch @@ -45,7 +47,7 @@ def test_registry_download_resolves_once_per_process() -> None: @pytest.fixture(autouse=True) -def _fresh_registry_cache(): +def _fresh_registry_cache() -> Iterator[None]: # registry_download memoizes per process; tests reuse package names registry.registry_download.cache_clear() yield @@ -112,7 +114,7 @@ def _http_response(text: str) -> MagicMock: return resp -def _registry_response(files: list[dict]): +def _registry_response(files: list[dict]) -> AbstractContextManager[MagicMock]: """Patch the consolidated HTTP path to serve a canned registry response.""" payload = {"versions": [{"name": "1.0.0", "files": files}]} return patch.object( @@ -308,7 +310,11 @@ def test_install_package_downloads_via_registry(tmp_path: Path) -> None: "pkg", "1.0.0", dest, [], tmp_path / "dl", expect=("payload",) ) assert mock_download.call_args[0][0] == "http://x/pkg.tar.gz" - assert mock_download.call_args[1] == {"sha256": "abc123", "size": 42} + assert mock_download.call_args[1] == { + "sha256": "abc123", + "size": 42, + "progress": None, + } def test_install_package_downloads_pinned(tmp_path: Path) -> None: @@ -334,7 +340,11 @@ def test_install_package_downloads_pinned(tmp_path: Path) -> None: ) mock_registry.assert_not_called() assert mock_download.call_args[0][0] == "http://y/pinned.tar.gz" - assert mock_download.call_args[1] == {"sha256": "def456", "size": 7} + assert mock_download.call_args[1] == { + "sha256": "def456", + "size": 7, + "progress": None, + } def test_install_package_mirror_wins_over_pinned(tmp_path: Path) -> None: @@ -546,8 +556,10 @@ def test_registry_download_non_list_system_is_named() -> None: registry.registry_download("pkg", "1.0.0") -def _resolve_for(sizes: dict[str, int | None]): - def resolve(name: str, version: str): +def _resolve_for( + sizes: dict[str, int | None], +) -> Callable[[str, str], tuple[str, str, int | None]]: + def resolve(name: str, version: str) -> tuple[str, str, int | None]: size = sizes[name] if size == -1: raise EsphomeError("registry down") @@ -567,8 +579,8 @@ def test_prefetch_packages_downloads_pending_in_parallel(tmp_path: Path) -> None ): registry.prefetch_packages( [ - ("a", "1.0", tmp_path / "a", []), - ("b", "2.0", tmp_path / "b", []), + ("a", "1.0", tmp_path / "a", [], ()), + ("b", "2.0", tmp_path / "b", [], ()), ], tmp_path / "dl", ) @@ -595,8 +607,8 @@ def test_prefetch_packages_uses_pinned_download(tmp_path: Path) -> None: ): registry.prefetch_packages( [ - ("a", "1.0", tmp_path / "a", []), - ("b", "2.0", tmp_path / "b", []), + ("a", "1.0", tmp_path / "a", [], ("payload",)), + ("b", "2.0", tmp_path / "b", [], ("payload",)), ], tmp_path / "dl", {"b": lambda: registry.Download("http://y/b.tar.gz", "def456", 20)}, @@ -626,7 +638,7 @@ def test_prefetch_packages_skips_freshly_installed_dest(tmp_path: Path) -> None: registry, "registry_download", side_effect=_resolve_for({"a": 10}) ), ): - registry.prefetch_packages([("a", "1.0", dest, [])], tmp_path / "dl") + registry.prefetch_packages([("a", "1.0", dest, [], ())], tmp_path / "dl") mock_download.assert_not_called() @@ -666,7 +678,7 @@ def test_prefetch_packages_waits_with_the_holders_progress( ), ): registry.prefetch_packages( - [("a", "1.0", dest, []), ("b", "2.0", tmp_path / "b", [])], + [("a", "1.0", dest, [], ()), ("b", "2.0", tmp_path / "b", [], ())], tmp_path / "dl", ) assert ticks == [0, 3, 10, 10] @@ -687,7 +699,10 @@ def test_prefetch_packages_leaves_a_long_held_lock_to_its_holder( ), ): registry.prefetch_packages( - [("a", "1.0", tmp_path / "a", []), ("b", "2.0", tmp_path / "b", [])], + [ + ("a", "1.0", tmp_path / "a", [], ()), + ("b", "2.0", tmp_path / "b", [], ()), + ], tmp_path / "dl", ) mock_download.assert_not_called() @@ -713,8 +728,8 @@ def test_prefetch_packages_dedupes_duplicate_entries(tmp_path: Path) -> None: ): registry.prefetch_packages( [ - ("a", "1.0", tmp_path / "a", []), - ("a", "1.0", tmp_path / "a", []), + ("a", "1.0", tmp_path / "a", [], ()), + ("a", "1.0", tmp_path / "a", [], ()), ], tmp_path / "dl", ) @@ -735,8 +750,8 @@ def test_prefetch_packages_single_pending_skips(tmp_path: Path) -> None: ): registry.prefetch_packages( [ - ("a", "1.0", marker_dest, []), - ("b", "2.0", tmp_path / "b", []), + ("a", "1.0", marker_dest, [], ()), + ("b", "2.0", tmp_path / "b", [], ()), ], tmp_path / "dl", ) @@ -758,9 +773,9 @@ def test_prefetch_packages_mirror_and_sizeless_stay_sequential( ): registry.prefetch_packages( [ - ("a", "1.0", tmp_path / "a", ["http://mirror/{VERSION}"]), - ("b", "2.0", tmp_path / "b", []), - ("c", "3.0", tmp_path / "c", []), + ("a", "1.0", tmp_path / "a", ["http://mirror/{VERSION}"], ()), + ("b", "2.0", tmp_path / "b", [], ()), + ("c", "3.0", tmp_path / "c", [], ()), ], tmp_path / "dl", ) @@ -781,8 +796,8 @@ def test_prefetch_packages_resolve_failure_defers_to_install( ): registry.prefetch_packages( [ - ("a", "1.0", tmp_path / "a", []), - ("b", "2.0", tmp_path / "b", []), + ("a", "1.0", tmp_path / "a", [], ()), + ("b", "2.0", tmp_path / "b", [], ()), ], tmp_path / "dl", ) @@ -803,8 +818,8 @@ def test_prefetch_packages_complete_archive_skipped(tmp_path: Path) -> None: ): registry.prefetch_packages( [ - ("a", "1.0", tmp_path / "a", []), - ("b", "2.0", tmp_path / "b", []), + ("a", "1.0", tmp_path / "a", [], ()), + ("b", "2.0", tmp_path / "b", [], ()), ], dl, ) @@ -826,8 +841,8 @@ def test_prefetch_packages_download_failure_is_debug( ): registry.prefetch_packages( [ - ("a", "1.0", tmp_path / "a", []), - ("b", "2.0", tmp_path / "b", []), + ("a", "1.0", tmp_path / "a", [], ()), + ("b", "2.0", tmp_path / "b", [], ()), ], tmp_path / "dl", ) @@ -851,9 +866,225 @@ def test_prefetch_packages_unexpected_failure_warns( ): registry.prefetch_packages( [ - ("a", "1.0", tmp_path / "a", []), - ("b", "2.0", tmp_path / "b", []), + ("a", "1.0", tmp_path / "a", [], ()), + ("b", "2.0", tmp_path / "b", [], ()), ], tmp_path / "dl", ) assert "TypeError" in caplog.text + + +def _spec( + name: str, + version: str, + dest: Path, + mirrors: list[str] | None = None, + expect: tuple[str, ...] = ("payload",), +) -> registry.PackageSpec: + return registry.PackageSpec(name, version, dest, mirrors or [], expect) + + +def test_install_packages_extracts_verified_archives_in_parallel( + tmp_path: Path, +) -> None: + """Two prefetched archives install concurrently under one shared bar.""" + dl = tmp_path / "dl" + dl.mkdir() + (dl / "a-1.0").write_bytes(b"x" * 10) + (dl / "b-2.0").write_bytes(b"y" * 20) + with patch.object(registry, "install_package") as mock_install: + registry.install_packages( + [_spec("a", "1.0", tmp_path / "a"), _spec("b", "2.0", tmp_path / "b")], dl + ) + assert mock_install.call_count == 2 + calls = sorted(mock_install.call_args_list, key=lambda c: c[0][0]) + for c, (name, version) in zip(calls, [("a", "1.0"), ("b", "2.0")], strict=True): + assert c[0][:3] == (name, version, tmp_path / name) + assert c[1]["expect"] == ("payload",) + assert callable(c[1]["extract_progress"]) + # Driving the tracker exercises the fraction-to-bytes scaling + c[1]["extract_progress"](0.5) + c[1]["extract_progress"](1.0) + + +def test_install_packages_single_archive_stays_sequential(tmp_path: Path) -> None: + """One verified archive has nothing to parallelize; original order kept.""" + dl = tmp_path / "dl" + dl.mkdir() + (dl / "a-1.0").write_bytes(b"x") + specs = [_spec("a", "1.0", tmp_path / "a"), _spec("b", "2.0", tmp_path / "b")] + with patch.object(registry, "install_package") as mock_install: + registry.install_packages(specs, dl) + assert [c[0][0] for c in mock_install.call_args_list] == ["a", "b"] + for c in mock_install.call_args_list: + assert "extract_progress" not in c[1] + + +def test_batched_download_progress_announces_a_real_download_once( + caplog: pytest.LogCaptureFixture, +) -> None: + """A batched archive that fails verification streams again behind a bar + that cannot move, so it says so once; a verified archive credits itself + in one tick and stays quiet.""" + ticks: list[float] = [] + with caplog.at_level(logging.INFO): + progress = registry._batched_download_progress("pkg", "1.0.0", ticks.append) + progress(0) + progress(4096) + assert caplog.text.count("Re-downloading pkg 1.0.0") == 1 + # The shared bar never moves for a download; it tracks extraction + assert ticks == [0.0, 0.0] + + # A resumed .part starts mid-file, so the first tick is not zero + caplog.clear() + ticks.clear() + with caplog.at_level(logging.INFO): + resumed = registry._batched_download_progress("pkg", "1.0.0", ticks.append) + resumed(8192) + resumed(16384) + assert caplog.text.count("Re-downloading pkg 1.0.0") == 1 + + caplog.clear() + ticks.clear() + with caplog.at_level(logging.INFO): + verified = registry._batched_download_progress("pkg", "1.0.0", ticks.append) + verified(42) + assert "Re-downloading" not in caplog.text + assert ticks == [0.0] + + +def test_install_packages_no_batch_logs_no_header( + tmp_path: Path, caplog: pytest.LogCaptureFixture +) -> None: + """The batch header must not describe a batch that never ran.""" + dl = tmp_path / "dl" + dl.mkdir() + specs = [_spec("a", "1.0", tmp_path / "a")] + with ( + caplog.at_level(logging.INFO), + patch.object(registry, "install_package"), + ): + registry.install_packages(specs, dl) + assert "Extracting 0" not in caplog.text + assert "package archive(s) with" not in caplog.text + + +def test_install_packages_mirror_and_marker_stay_sequential(tmp_path: Path) -> None: + """Mirror overrides and marker hits never enter the parallel batch.""" + dl = tmp_path / "dl" + dl.mkdir() + for name, ver in (("a", "1.0"), ("b", "2.0"), ("c", "3.0"), ("d", "4.0")): + (dl / f"{name}-{ver}").write_bytes(b"x") + marked = tmp_path / "c" + marked.mkdir() + (marked / ".esphome_extracted").touch() + specs = [ + _spec("a", "1.0", tmp_path / "a"), + _spec("b", "2.0", tmp_path / "b", mirrors=["http://m"]), + _spec("c", "3.0", marked), + _spec("d", "4.0", tmp_path / "d"), + ] + with patch.object(registry, "install_package") as mock_install: + registry.install_packages(specs, dl) + sequential = [ + c for c in mock_install.call_args_list if "extract_progress" not in c[1] + ] + batched = [c for c in mock_install.call_args_list if "extract_progress" in c[1]] + assert sorted(c[0][0] for c in sequential) == ["b", "c"] + assert sorted(c[0][0] for c in batched) == ["a", "d"] + + +def test_install_packages_first_failure_reraised( + tmp_path: Path, caplog: pytest.LogCaptureFixture +) -> None: + """Installs are mandatory: the first failure propagates, extras are logged.""" + dl = tmp_path / "dl" + dl.mkdir() + (dl / "a-1.0").write_bytes(b"x") + (dl / "b-2.0").write_bytes(b"y") + boom = EsphomeError("bad layout") + + def _fail(name: str, *_a, **_kw) -> None: + raise boom if name == "a" else EsphomeError("also bad") + + with ( + patch.object(registry, "install_package", side_effect=_fail), + pytest.raises(EsphomeError), + ): + registry.install_packages( + [_spec("a", "1.0", tmp_path / "a"), _spec("b", "2.0", tmp_path / "b")], dl + ) + # Every failure is named, including the re-raised one: its exception + # message may not identify the package + assert "Could not install a" in caplog.text + assert "Could not install b" in caplog.text + + +@contextmanager +def _batched_install( + tmp_path: Path, + extract_progress: Callable[[float], None] | None, + prefill_archive: bool = True, +) -> Iterator[tuple[MagicMock, MagicMock]]: + """Run a batched install_package of pkg@1.0.0; yields the download and + extract mocks.""" + dest = tmp_path / "pkg" + if prefill_archive: + (tmp_path / "dl").mkdir() + (tmp_path / "dl" / "pkg-1.0.0").write_bytes(b"x") + with ( + patch.object(registry, "download_with_resume") as mock_download, + patch.object(registry, "archive_extract_all") as mock_extract, + patch.object( + registry, + "registry_download", + return_value=("http://x/pkg.tar.gz", "abc123", 42), + ), + ): + mock_extract.side_effect = lambda *_a, **_kw: (dest / "payload").mkdir( + parents=True + ) + registry.install_package( + "pkg", + "1.0.0", + dest, + [], + tmp_path / "dl", + expect=("payload",), + extract_progress=extract_progress, + ) + yield mock_download, mock_extract + + +def test_install_package_extract_progress_suppresses_bars( + tmp_path: Path, caplog: pytest.LogCaptureFixture +) -> None: + """A batched install routes extraction fractions to the caller and keeps + both private bars and per-package INFO lines off the shared bar.""" + fractions: list[float] = [] + with ( + caplog.at_level(logging.INFO), + _batched_install(tmp_path, fractions.append) as (mock_download, mock_extract), + ): + pass + assert mock_extract.call_args[1]["progress"] == fractions.append + # The download tracker reports zero bytes, keeping the shared bar honest + download_progress = mock_download.call_args[1]["progress"] + assert callable(download_progress) + download_progress(42) + assert fractions == [0.0] + assert "Downloading pkg" not in caplog.text + assert "Extracting pkg" not in caplog.text + + +def test_install_package_batched_missing_archive_keeps_info_log( + tmp_path: Path, caplog: pytest.LogCaptureFixture +) -> None: + """A batched archive that unexpectedly needs a real download keeps the + INFO line; the shared bar shows no progress for it.""" + with ( + caplog.at_level(logging.INFO), + _batched_install(tmp_path, lambda _frac: None, prefill_archive=False), + ): + pass + assert "Downloading pkg 1.0.0" in caplog.text