From 0a8e4c8db714a3f91c9de842f648baf631655685 Mon Sep 17 00:00:00 2001 From: Paulus Schoutsen Date: Thu, 24 Sep 2026 11:47:43 +0100 Subject: [PATCH] Run mypy once for all enum_identity_compare cases (#180461) Co-authored-by: Franck Nijhof Co-authored-by: Markus Tuominen <3738613+Markus98@users.noreply.github.com> Co-authored-by: Claude Opus 5 --- .../test_enum_identity_compare.py | 650 ++++++++++-------- 1 file changed, 346 insertions(+), 304 deletions(-) diff --git a/tests/mypy_plugins/test_enum_identity_compare.py b/tests/mypy_plugins/test_enum_identity_compare.py index 1f7b34152323..c3a84f8c6272 100644 --- a/tests/mypy_plugins/test_enum_identity_compare.py +++ b/tests/mypy_plugins/test_enum_identity_compare.py @@ -1,8 +1,9 @@ """Tests for the enum_identity_compare mypy plugin. -Each test snippet is run through mypy in a subprocess with the plugin enabled. -Tests assert the number of ``home-assistant-enum-identity-compare`` errors emitted -and the relevant message content (operator pair and enum class name). +Every test snippet is type-checked by a single mypy subprocess with the plugin +enabled, and each test asserts on the errors reported for its own snippet: the +number of ``home-assistant-enum-identity-compare`` errors emitted and the +relevant message content (operator pair and enum class name). The plugin is intentionally narrow: it fires only on plain ``enum.Enum`` subclasses (where ``__eq__`` is identity-based) plus a small set of @@ -25,26 +26,37 @@ _PROJECT_ROOT = Path(__file__).resolve().parents[2] _PLUGINS_ROOT = _PROJECT_ROOT # mypy_plugins/ lives under the worktree root -def _run_mypy(code: str, tmp_path: Path, mypy_path: str | None = None) -> list[str]: - """Run mypy with the plugin and return home-assistant-enum-identity-compare errors. +def _write_flow_result_type_stub(tmp_path: Path) -> None: + """Write a fake ``homeassistant.data_entry_flow`` module under tmp_path. - Each error is normalized to ``LINE: MESSAGE`` form. ``mypy_path``, if - given, is written into ``mypy.ini`` so tests can supply stub modules - that resolve to specific fullnames (used for the framework-guaranteed set). + ``_FRAMEWORK_GUARANTEED_ENUMS`` matches by fullname + (``homeassistant.data_entry_flow.FlowResultType``), so we synthesize a + package at that path rather than depending on the real HA tree being + importable from the test environment. """ - src = tmp_path / "case.py" - src.write_text(textwrap.dedent(code)) - cache = tmp_path / "mypy_cache" + pkg = tmp_path / "homeassistant" + pkg.mkdir() + (pkg / "__init__.py").write_text("") + (pkg / "data_entry_flow.py").write_text( + "from enum import StrEnum\n" + "class FlowResultType(StrEnum):\n" + ' FORM = "form"\n' + ' ABORT = "abort"\n' + ) + + +def _run_mypy(paths: list[Path], tmp_path: Path) -> str: + """Run mypy with the plugin over paths and return its stdout.""" config = tmp_path / "mypy.ini" - config_body = ( + config.write_text( "[mypy]\n" "plugins = mypy_plugins.enum_identity_compare\n" "show_error_codes = true\n" "strict_equality = true\n" + # Lets the stub package written next to the cases resolve to the + # fullname the framework-guaranteed set matches on. + f"mypy_path = {tmp_path}\n" ) - if mypy_path is not None: - config_body += f"mypy_path = {mypy_path}\n" - config.write_text(config_body) # Run mypy in a subprocess: type checking a snippet in process leaves # hundreds of thousands of objects behind, which the next test module pays @@ -53,16 +65,16 @@ def _run_mypy(code: str, tmp_path: Path, mypy_path: str | None = None) -> list[s env["PYTHONPATH"] = os.pathsep.join( filter(None, (str(_PLUGINS_ROOT), env.get("PYTHONPATH"))) ) - stdout = subprocess.run( + return subprocess.run( [ sys.executable, "-m", "mypy", "--no-incremental", - f"--cache-dir={cache}", + f"--cache-dir={tmp_path / 'mypy_cache'}", "--config-file", str(config), - str(src), + *(str(path) for path in paths), ], capture_output=True, check=False, @@ -70,16 +82,252 @@ def _run_mypy(code: str, tmp_path: Path, mypy_path: str | None = None) -> list[s env=env, ).stdout - errors: list[str] = [] - for line in stdout.splitlines(): - if "[home-assistant-enum-identity-compare]" not in line: - continue - # Format: ":: error: [code]" - prefix, _, msg = line.partition(": error: ") - line_no = prefix.rsplit(":", 1)[-1].strip() - msg_clean = msg.split(" [home-assistant-enum-identity-compare]", 1)[0].strip() - errors.append(f"{line_no}: {msg_clean}") - return errors + +_PLAIN_ENUM_CASES = [ + pytest.param( + """ +def fn(s: ConfigEntryState) -> bool: + return s == ConfigEntryState.LOADED +""", + "ConfigEntryState", + _IS_EQ, + id="plain_enum_eq", + ), + pytest.param( + """ +def fn(s: ConfigEntryState) -> bool: + return s != ConfigEntryState.NOT_LOADED +""", + "ConfigEntryState", + _IS_NOT_NE, + id="plain_enum_ne", + ), + pytest.param( + """ +def fn(s: ConfigEntryState) -> bool: + return ConfigEntryState.LOADED == s +""", + "ConfigEntryState", + _IS_EQ, + id="plain_enum_lhs", + ), + # An ``elif`` after an ``is`` check narrows the LHS to a literal + # union. Without union/literal handling, the plugin would silently + # skip the ``==`` even though both operands resolve to ``SourceCodes``. + pytest.param( + """ +def fn(source: SourceCodes) -> str: + if source is SourceCodes.DAB: + return "dab" + elif source == SourceCodes.FM: + return "fm" + return "other" +""", + "SourceCodes", + _IS_EQ, + id="narrowed_elif", + ), + pytest.param( + """ +from typing import Literal + +def fn(s: Literal[ConfigEntryState.LOADED]) -> bool: + return s == ConfigEntryState.LOADED +""", + "ConfigEntryState", + _IS_EQ, + id="literal_annotation", + ), + # ``Enum | None == Enum``: the plugin conservatively rejects any union + # containing ``None``, but under ``strict_equality`` mypy narrows the + # LHS to the enum class before invoking ``__eq__``, so this call site + # never reaches the union path and is correctly flagged. HA's + # ``mypy.ini`` sets ``strict_equality``. + pytest.param( + """ +def fn(source: SourceCodes | None) -> bool: + return source == SourceCodes.DAB +""", + "SourceCodes", + _IS_EQ, + id="optional_enum_under_strict_equality", + ), + # An enum deriving from an intermediate ``Enum`` base (no data mixin) + # is still identity-based and must be flagged — the structural check + # must not mistake the intermediate base for a value mixin. + pytest.param( + """ +def fn(s: DerivedStates) -> bool: + return s == DerivedStates.ON +""", + "DerivedStates", + _IS_EQ, + id="derived_from_intermediate_enum_base", + ), +] + +_NO_FLAG_CASES = [ + # A StrEnum defined in user code is NOT flagged: the plugin can't tell + # whether callers pass the enum instance or the underlying string + # (StrEnum's whole point is making both work). It must be added to + # ``_FRAMEWORK_GUARANTEED_ENUMS`` to be checked. + pytest.param( + """ +def fn(m: MediaType) -> bool: + return m == MediaType.CHANNEL +""", + id="strenum_not_framework_guaranteed", + ), + # An IntEnum defined in user code is NOT flagged for the same reason. + pytest.param( + """ +def fn(code: HTTPStatus) -> bool: + return code == HTTPStatus.OK +""", + id="intenum_not_framework_guaranteed", + ), + # The legacy ``(int, Enum)`` mixin inherits a value-based ``__eq__`` + # from ``int`` (no ``enum.IntEnum`` base), so it is NOT flagged either. + pytest.param( + """ +def fn(rate: AudioBitRates) -> bool: + return rate == AudioBitRates.BITRATE_8 +""", + id="int_enum_mixin_not_framework_guaranteed", + ), + # And the same for the legacy ``(str, Enum)`` mixin. + pytest.param( + """ +def fn(v: LegacyStr) -> bool: + return v == LegacyStr.A +""", + id="str_enum_mixin_not_framework_guaranteed", + ), + # A ``(float, Enum)`` mixin is value-based too (``__eq__`` from float). + pytest.param( + """ +def fn(s: HomeeCoverState) -> bool: + return s == HomeeCoverState.OPEN +""", + id="float_enum_mixin_not_framework_guaranteed", + ), + # And a ``(bytes, Enum)`` mixin likewise. + pytest.param( + """ +def fn(v: LegacyBytes) -> bool: + return v == LegacyBytes.A +""", + id="bytes_enum_mixin_not_framework_guaranteed", + ), + # A ``@dataclass`` mixin is value-based (generated ``__eq__`` compares + # by value), even though the mixin is not a builtin primitive. + pytest.param( + """ +def fn(v: HardwareVariant) -> bool: + return v == HardwareVariant.A +""", + id="dataclass_mixin_not_flagged", + ), + # A ``NamedTuple`` mixin is value-based too (tuple ``__eq__``). + pytest.param( + """ +def fn(v: NamedTupleEnum) -> bool: + return v == NamedTupleEnum.A +""", + id="namedtuple_mixin_not_flagged", + ), + # Comparing a raw ``str`` against a ``StrEnum`` member is legitimate. + pytest.param( + """ +def fn(raw: str) -> bool: + return raw == MediaType.CHANNEL +""", + id="str_vs_strenum", + ), + # Comparing a raw ``int`` against an ``IntEnum`` member is legitimate. + pytest.param( + """ +def fn(code: int) -> bool: + return code == HTTPStatus.OK +""", + id="int_vs_intenum", + ), + # A union LHS (e.g. ``MediaType | str``) is NOT flagged: even if + # MediaType were framework-guaranteed, runtime callers can pass either form, so + # switching to ``is`` would break the str arm. + pytest.param( + """ +def fn(m: MediaType | str) -> bool: + return m == MediaType.CHANNEL +""", + id="union_with_str", + ), + # ``IntFlag`` bitwise ``==`` is the standard pattern. + pytest.param( + """ +def fn(features: ClimateFeature) -> bool: + return features & ClimateFeature.SWING_MODE == ClimateFeature.SWING_MODE +""", + id="intflag_bitwise", + ), + # ``is`` is the recommended form — must not fire on itself. + pytest.param( + """ +def fn(s: ConfigEntryState) -> bool: + return s is ConfigEntryState.LOADED +""", + id="is_already", + ), + # ``is not`` is the recommended negative form — must not fire on itself. + pytest.param( + """ +def fn(s: ConfigEntryState) -> bool: + return s is not ConfigEntryState.LOADED +""", + id="is_not_already", + ), + # Plain ``int`` ``==`` ``int`` (no enum involved) must not flag. + pytest.param( + "\nif 1 == 2:\n pass\n", + id="unrelated_compare", + ), + # Ordering comparison (``>``) on an enum must not flag. + pytest.param( + """ +def fn(s: HTTPStatus) -> bool: + return s > HTTPStatus.OK +""", + id="ordering_op", + ), +] + +_FRAMEWORK_CASES = [ + # ``FlowResultType`` is on ``_FRAMEWORK_GUARANTEED_ENUMS`` and must + # flag. StrEnum normally escapes the plugin, but this class is + # explicitly included because HA's framework controls every + # value-assigning callsite — see the plugin module docstring. + pytest.param( + """ +from homeassistant.data_entry_flow import FlowResultType + +def fn(r: FlowResultType) -> bool: + return r == FlowResultType.FORM +""", + _IS_EQ, + id="eq", + ), + # And the same for ``!=`` against the framework-guaranteed ``FlowResultType``. + pytest.param( + """ +from homeassistant.data_entry_flow import FlowResultType + +def fn(r: FlowResultType) -> bool: + return r != FlowResultType.ABORT +""", + _IS_NOT_NE, + id="ne", + ), +] _PRELUDE = """ @@ -146,304 +394,98 @@ class DerivedStates(_BaseStates): """ +@pytest.fixture(scope="module") +def mypy_errors(tmp_path_factory: pytest.TempPathFactory) -> dict[str, list[str]]: + """Type-check every case in one mypy run and return the errors per case id. + + Analysing typeshed dominates a mypy run and dwarfs the few lines each case + contributes. The cases never import each other, so checking them together + reports the same errors as checking them one at a time. + + Cases are keyed by their ``pytest.param`` id, which must be unique across + all three tables, so two cases that happen to share a snippet stay separate. + + Each error is normalized to ``LINE: MESSAGE`` form. + """ + tmp_path = tmp_path_factory.mktemp("enum_identity_compare") + _write_flow_result_type_stub(tmp_path) + + sources = { + **{ + case.id: _PRELUDE + case.values[0] + for case in (*_PLAIN_ENUM_CASES, *_NO_FLAG_CASES) + }, + **{case.id: case.values[0] for case in _FRAMEWORK_CASES}, + } + assert len(sources) == len(_PLAIN_ENUM_CASES) + len(_NO_FLAG_CASES) + len( + _FRAMEWORK_CASES + ), "Duplicate pytest.param id across case tables" + + case_id_by_module: dict[str, str] = {} + paths: list[Path] = [] + for index, (case_id, source) in enumerate(sources.items()): + path = tmp_path / f"case_{index}.py" + path.write_text(textwrap.dedent(source)) + case_id_by_module[path.stem] = case_id + paths.append(path) + + errors: dict[str, list[str]] = {case_id: [] for case_id in sources} + for line in _run_mypy(paths, tmp_path).splitlines(): + if "[home-assistant-enum-identity-compare]" not in line: + continue + # Format: ":: error: [code]" + prefix, _, msg = line.partition(": error: ") + path_text, _, line_no = prefix.rpartition(":") + msg_clean = msg.split(" [home-assistant-enum-identity-compare]", 1)[0].strip() + errors[case_id_by_module[Path(path_text).stem]].append( + f"{line_no.strip()}: {msg_clean}" + ) + return errors + + +@pytest.fixture +def case_errors( + request: pytest.FixtureRequest, mypy_errors: dict[str, list[str]] +) -> list[str]: + """Return the errors mypy reported for the running case.""" + return mypy_errors[request.node.callspec.id] + + @pytest.mark.parametrize( ("snippet", "enum_name", "op_substrings"), - [ - pytest.param( - """ -def fn(s: ConfigEntryState) -> bool: - return s == ConfigEntryState.LOADED -""", - "ConfigEntryState", - _IS_EQ, - id="plain_enum_eq", - ), - pytest.param( - """ -def fn(s: ConfigEntryState) -> bool: - return s != ConfigEntryState.NOT_LOADED -""", - "ConfigEntryState", - _IS_NOT_NE, - id="plain_enum_ne", - ), - pytest.param( - """ -def fn(s: ConfigEntryState) -> bool: - return ConfigEntryState.LOADED == s -""", - "ConfigEntryState", - _IS_EQ, - id="plain_enum_lhs", - ), - # An ``elif`` after an ``is`` check narrows the LHS to a literal - # union. Without union/literal handling, the plugin would silently - # skip the ``==`` even though both operands resolve to ``SourceCodes``. - pytest.param( - """ -def fn(source: SourceCodes) -> str: - if source is SourceCodes.DAB: - return "dab" - elif source == SourceCodes.FM: - return "fm" - return "other" -""", - "SourceCodes", - _IS_EQ, - id="narrowed_elif", - ), - pytest.param( - """ -from typing import Literal - -def fn(s: Literal[ConfigEntryState.LOADED]) -> bool: - return s == ConfigEntryState.LOADED -""", - "ConfigEntryState", - _IS_EQ, - id="literal_annotation", - ), - # ``Enum | None == Enum``: the plugin conservatively rejects any union - # containing ``None``, but under ``strict_equality`` mypy narrows the - # LHS to the enum class before invoking ``__eq__``, so this call site - # never reaches the union path and is correctly flagged. HA's - # ``mypy.ini`` sets ``strict_equality``. - pytest.param( - """ -def fn(source: SourceCodes | None) -> bool: - return source == SourceCodes.DAB -""", - "SourceCodes", - _IS_EQ, - id="optional_enum_under_strict_equality", - ), - # An enum deriving from an intermediate ``Enum`` base (no data mixin) - # is still identity-based and must be flagged — the structural check - # must not mistake the intermediate base for a value mixin. - pytest.param( - """ -def fn(s: DerivedStates) -> bool: - return s == DerivedStates.ON -""", - "DerivedStates", - _IS_EQ, - id="derived_from_intermediate_enum_base", - ), - ], + _PLAIN_ENUM_CASES, ) def test_bad_plain_enum( - tmp_path: Path, + case_errors: list[str], snippet: str, enum_name: str, op_substrings: tuple[str, str], ) -> None: """Comparisons on plain ``Enum`` operands must flag a single error.""" - errors = _run_mypy(_PRELUDE + snippet, tmp_path) - assert len(errors) == 1 - assert enum_name in errors[0] - assert all(op in errors[0] for op in op_substrings) + assert len(case_errors) == 1 + assert enum_name in case_errors[0] + assert all(op in case_errors[0] for op in op_substrings) @pytest.mark.parametrize( "snippet", - [ - # A StrEnum defined in user code is NOT flagged: the plugin can't tell - # whether callers pass the enum instance or the underlying string - # (StrEnum's whole point is making both work). It must be added to - # ``_FRAMEWORK_GUARANTEED_ENUMS`` to be checked. - pytest.param( - """ -def fn(m: MediaType) -> bool: - return m == MediaType.CHANNEL -""", - id="strenum_not_framework_guaranteed", - ), - # An IntEnum defined in user code is NOT flagged for the same reason. - pytest.param( - """ -def fn(code: HTTPStatus) -> bool: - return code == HTTPStatus.OK -""", - id="intenum_not_framework_guaranteed", - ), - # The legacy ``(int, Enum)`` mixin inherits a value-based ``__eq__`` - # from ``int`` (no ``enum.IntEnum`` base), so it is NOT flagged either. - pytest.param( - """ -def fn(rate: AudioBitRates) -> bool: - return rate == AudioBitRates.BITRATE_8 -""", - id="int_enum_mixin_not_framework_guaranteed", - ), - # And the same for the legacy ``(str, Enum)`` mixin. - pytest.param( - """ -def fn(v: LegacyStr) -> bool: - return v == LegacyStr.A -""", - id="str_enum_mixin_not_framework_guaranteed", - ), - # A ``(float, Enum)`` mixin is value-based too (``__eq__`` from float). - pytest.param( - """ -def fn(s: HomeeCoverState) -> bool: - return s == HomeeCoverState.OPEN -""", - id="float_enum_mixin_not_framework_guaranteed", - ), - # And a ``(bytes, Enum)`` mixin likewise. - pytest.param( - """ -def fn(v: LegacyBytes) -> bool: - return v == LegacyBytes.A -""", - id="bytes_enum_mixin_not_framework_guaranteed", - ), - # A ``@dataclass`` mixin is value-based (generated ``__eq__`` compares - # by value), even though the mixin is not a builtin primitive. - pytest.param( - """ -def fn(v: HardwareVariant) -> bool: - return v == HardwareVariant.A -""", - id="dataclass_mixin_not_flagged", - ), - # A ``NamedTuple`` mixin is value-based too (tuple ``__eq__``). - pytest.param( - """ -def fn(v: NamedTupleEnum) -> bool: - return v == NamedTupleEnum.A -""", - id="namedtuple_mixin_not_flagged", - ), - # Comparing a raw ``str`` against a ``StrEnum`` member is legitimate. - pytest.param( - """ -def fn(raw: str) -> bool: - return raw == MediaType.CHANNEL -""", - id="str_vs_strenum", - ), - # Comparing a raw ``int`` against an ``IntEnum`` member is legitimate. - pytest.param( - """ -def fn(code: int) -> bool: - return code == HTTPStatus.OK -""", - id="int_vs_intenum", - ), - # A union LHS (e.g. ``MediaType | str``) is NOT flagged: even if - # MediaType were framework-guaranteed, runtime callers can pass either form, so - # switching to ``is`` would break the str arm. - pytest.param( - """ -def fn(m: MediaType | str) -> bool: - return m == MediaType.CHANNEL -""", - id="union_with_str", - ), - # ``IntFlag`` bitwise ``==`` is the standard pattern. - pytest.param( - """ -def fn(features: ClimateFeature) -> bool: - return features & ClimateFeature.SWING_MODE == ClimateFeature.SWING_MODE -""", - id="intflag_bitwise", - ), - # ``is`` is the recommended form — must not fire on itself. - pytest.param( - """ -def fn(s: ConfigEntryState) -> bool: - return s is ConfigEntryState.LOADED -""", - id="is_already", - ), - # ``is not`` is the recommended negative form — must not fire on itself. - pytest.param( - """ -def fn(s: ConfigEntryState) -> bool: - return s is not ConfigEntryState.LOADED -""", - id="is_not_already", - ), - # Plain ``int`` ``==`` ``int`` (no enum involved) must not flag. - pytest.param( - "\nif 1 == 2:\n pass\n", - id="unrelated_compare", - ), - # Ordering comparison (``>``) on an enum must not flag. - pytest.param( - """ -def fn(s: HTTPStatus) -> bool: - return s > HTTPStatus.OK -""", - id="ordering_op", - ), - ], + _NO_FLAG_CASES, ) -def test_good_no_flag(tmp_path: Path, snippet: str) -> None: +def test_good_no_flag(case_errors: list[str], snippet: str) -> None: """Legitimate comparisons must not emit any error.""" - errors = _run_mypy(_PRELUDE + snippet, tmp_path) - assert errors == [] - - -def _write_flow_result_type_stub(tmp_path: Path) -> None: - """Write a fake ``homeassistant.data_entry_flow`` module under tmp_path. - - ``_FRAMEWORK_GUARANTEED_ENUMS`` matches by fullname - (``homeassistant.data_entry_flow.FlowResultType``), so we synthesize a - package at that path rather than depending on the real HA tree being - importable from the test environment. - """ - pkg = tmp_path / "homeassistant" - pkg.mkdir() - (pkg / "__init__.py").write_text("") - (pkg / "data_entry_flow.py").write_text( - "from enum import StrEnum\n" - "class FlowResultType(StrEnum):\n" - ' FORM = "form"\n' - ' ABORT = "abort"\n' - ) + assert case_errors == [] @pytest.mark.parametrize( ("snippet", "op_substrings"), - [ - # ``FlowResultType`` is on ``_FRAMEWORK_GUARANTEED_ENUMS`` and must - # flag. StrEnum normally escapes the plugin, but this class is - # explicitly included because HA's framework controls every - # value-assigning callsite — see the plugin module docstring. - pytest.param( - """ -from homeassistant.data_entry_flow import FlowResultType - -def fn(r: FlowResultType) -> bool: - return r == FlowResultType.FORM -""", - _IS_EQ, - id="eq", - ), - # And the same for ``!=`` against the framework-guaranteed ``FlowResultType``. - pytest.param( - """ -from homeassistant.data_entry_flow import FlowResultType - -def fn(r: FlowResultType) -> bool: - return r != FlowResultType.ABORT -""", - _IS_NOT_NE, - id="ne", - ), - ], + _FRAMEWORK_CASES, ) def test_bad_framework_guaranteed( - tmp_path: Path, + case_errors: list[str], snippet: str, op_substrings: tuple[str, str], ) -> None: """The framework-guaranteed ``FlowResultType`` StrEnum must flag a single error.""" - _write_flow_result_type_stub(tmp_path) - errors = _run_mypy(snippet, tmp_path, mypy_path=str(tmp_path)) - assert len(errors) == 1 - assert "FlowResultType" in errors[0] - assert all(op in errors[0] for op in op_substrings) + assert len(case_errors) == 1 + assert "FlowResultType" in case_errors[0] + assert all(op in case_errors[0] for op in op_substrings)