diff --git a/esphome/components/api/wizard.py b/esphome/components/api/wizard.py index 314b907e4d..befed39972 100644 --- a/esphome/components/api/wizard.py +++ b/esphome/components/api/wizard.py @@ -6,7 +6,6 @@ Home Assistant shows the wizard when the device is added. See api_wizard.h for t from collections.abc import Callable import importlib import json -from types import ModuleType from typing import Any from esphome import automation @@ -27,7 +26,12 @@ from esphome.const import ( ) from esphome.core import CORE, ID import esphome.final_validate as fv -from esphome.helpers import fnv1_hash, fnv1_hash_object_id, fnv1a_32bit_hash +from esphome.helpers import ( + fnv1_hash, + fnv1_hash_object_id, + fnv1a_32bit_hash, + zstd_module, +) from esphome.types import ConfigType API_DOMAIN = "api" @@ -351,14 +355,6 @@ automation.register_apply_condition( ) -def zstd_module() -> ModuleType: - """The zstd module: the standard library one from Python 3.14, otherwise the backport.""" - try: - return importlib.import_module("compression.zstd") - except ImportError: - return importlib.import_module("backports.zstd") - - def _entity_document(conf: ConfigType, config: fv.FinalValidateConfig) -> ConfigType: """An entity of the device that the page shows, keyed as ListEntitiesResponse keys it.""" _, declaration = _wizard_input_declaration(config, conf[CONF_ID]) diff --git a/esphome/helpers.py b/esphome/helpers.py index 5ec46be67e..1f9b140797 100644 --- a/esphome/helpers.py +++ b/esphome/helpers.py @@ -2,6 +2,7 @@ from __future__ import annotations from collections.abc import Callable, Iterable, MutableMapping from contextlib import suppress +import importlib import ipaddress import logging import os @@ -10,6 +11,7 @@ import platform import re import stat import sys +from types import ModuleType from typing import TYPE_CHECKING, TextIO from esphome.const import __version__ as ESPHOME_VERSION @@ -827,3 +829,11 @@ def docs_url(path: str) -> str: path = path.removeprefix("/") return docs_format.format(path=path) + + +def zstd_module() -> ModuleType: + """The zstd module: the standard library one from Python 3.14, otherwise the backport.""" + try: + return importlib.import_module("compression.zstd") + except ImportError: + return importlib.import_module("backports.zstd") diff --git a/requirements.txt b/requirements.txt index 6e1dfa1ba5..8deafbad86 100644 --- a/requirements.txt +++ b/requirements.txt @@ -19,7 +19,7 @@ puremagic==2.2.0 ruamel.yaml==0.19.1 # dashboard_import ruamel.yaml.clib==0.2.15 # dashboard_import esphome-glyphsets==0.2.0 -backports.zstd==1.7.0; python_version < "3.14" # api wizard; compression.zstd is in the standard library from 3.14 +backports.zstd==1.7.0; python_version < "3.14" # esphome.helpers.zstd_module; compression.zstd is in the standard library from 3.14 pillow==12.3.0 resvg-py==0.5.0 freetype-py==2.5.1 diff --git a/tests/component_tests/api/test_wizard.py b/tests/component_tests/api/test_wizard.py index 276f42d624..112092482b 100644 --- a/tests/component_tests/api/test_wizard.py +++ b/tests/component_tests/api/test_wizard.py @@ -14,7 +14,7 @@ from esphome.components.api import wizard from esphome.components.homeassistant.switch import SUPPORTED_DOMAINS as SWITCH_DOMAINS from esphome.config import load_config from esphome.core import CORE -from esphome.helpers import fnv1_hash, fnv1a_32bit_hash +from esphome.helpers import fnv1_hash, fnv1a_32bit_hash, zstd_module from tests.component_tests.helpers import get_define_value ESP32_HEADER = """ @@ -613,7 +613,7 @@ def test_text_select_and_button_default_to_their_domains( tmp_path: Path, generate_main: Callable[[str | Path], str] ) -> None: main_cpp = generate_main(write_input_config(tmp_path, ESP32_HEADER)) - document = json.loads(wizard.zstd_module().decompress(blob_in(main_cpp))) + document = json.loads(zstd_module().decompress(blob_in(main_cpp))) filters = { entry["key"]: entry.get("entity_filters") for page in document["pages"] @@ -733,7 +733,7 @@ def test_the_blob_is_the_exact_json_document( ) blob = blob_in(main_cpp) - text = wizard.zstd_module().decompress(blob).decode("utf-8") + text = zstd_module().decompress(blob).decode("utf-8") expected = { "version": 1, "pages": [ @@ -788,7 +788,7 @@ def test_the_blob_is_deterministic_and_the_same_on_every_platform( ) # Compressing the same document again gives the same bytes document = wizard.wizard_document(CORE.config["api"]["wizard"], CORE.config) - again = wizard.zstd_module().compress( + again = zstd_module().compress( json.dumps( document, separators=(",", ":"), sort_keys=True, ensure_ascii=False ).encode("utf-8"), @@ -859,21 +859,6 @@ def test_an_entity_without_a_name_cannot_be_in_the_wizard(tmp_path: Path) -> Non assert any("has no name of its own" in error for error in errors), errors -def test_zstd_falls_back_to_the_backport(monkeypatch: pytest.MonkeyPatch) -> None: - """Before Python 3.14 the standard library has no zstd, so the backport is used.""" - backport = object() - - def import_module(name: str) -> object: - if name == "compression.zstd": - raise ImportError(name) - assert name == "backports.zstd" - return backport - - monkeypatch.setattr(wizard.importlib, "import_module", import_module) - - assert wizard.zstd_module() is backport - - def test_a_wizard_too_big_for_one_message_is_rejected(tmp_path: Path) -> None: # Random text does not compress, so this needs more than the limit even compressed rng = random.Random(1) diff --git a/tests/unit_tests/test_helpers.py b/tests/unit_tests/test_helpers.py index 8dae1e87b2..764f160cbc 100644 --- a/tests/unit_tests/test_helpers.py +++ b/tests/unit_tests/test_helpers.py @@ -1285,3 +1285,20 @@ def test_get_usable_cpu_count_sources() -> None: mock_os_unknown = types.SimpleNamespace(cpu_count=lambda: None) with patch("esphome.helpers.os", mock_os_unknown): assert helpers.get_usable_cpu_count() == 1 + + +def test_zstd_module_falls_back_to_the_backport( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Before Python 3.14 the standard library has no zstd, so the backport is used.""" + backport = object() + + def import_module(name: str) -> object: + if name == "compression.zstd": + raise ImportError(name) + assert name == "backports.zstd" + return backport + + monkeypatch.setattr(helpers.importlib, "import_module", import_module) + + assert helpers.zstd_module() is backport