[core] Extract native toolchain package archives in parallel (#18840)

This commit is contained in:
J. Nick Koston
2026-09-30 02:23:05 +02:00
committed by GitHub
parent ef0819f983
commit ed88002c37
10 changed files with 591 additions and 137 deletions
+8 -15
View File
@@ -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
)
+2 -2
View File
@@ -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
+70 -30
View File
@@ -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(
+2 -2
View File
@@ -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
+5 -4
View File
@@ -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.
+140 -15
View File
@@ -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),
)
+16 -35
View File
@@ -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:
+87 -4
View File
@@ -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"),
[
+1 -1
View File
@@ -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,
[
+260 -29
View File
@@ -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