Add RepairsFlowContext to RepairsFlow (#180942)

This commit is contained in:
iluvdata
2026-09-11 01:44:16 +02:00
committed by GitHub
parent cc09d56065
commit d5b6bd957d
6 changed files with 201 additions and 18 deletions
+2 -1
View File
@@ -7,13 +7,14 @@ from homeassistant.helpers.typing import ConfigType
from . import issue_handler, websocket_api
from .const import DOMAIN, FlowType
from .issue_handler import ConfirmRepairFlow, RepairsFlowManager
from .models import RepairsFlow, RepairsFlowResult
from .models import RepairsFlow, RepairsFlowContext, RepairsFlowResult
__all__ = [
"DOMAIN",
"ConfirmRepairFlow",
"FlowType",
"RepairsFlow",
"RepairsFlowContext",
"RepairsFlowManager",
"RepairsFlowResult",
"repairs_flow_manager",
@@ -11,7 +11,7 @@ from homeassistant.helpers import issue_registry as ir
from homeassistant.helpers.integration_platform import LazyIntegrationPlatforms
from .const import DOMAIN
from .models import RepairsFlow, RepairsFlowResult, RepairsProtocol
from .models import RepairsFlow, RepairsFlowContext, RepairsFlowResult, RepairsProtocol
class ConfirmRepairFlow(RepairsFlow):
@@ -43,21 +43,41 @@ class ConfirmRepairFlow(RepairsFlow):
class RepairsFlowManager(
data_entry_flow.FlowManager[data_entry_flow.FlowContext, RepairsFlowResult, str]
data_entry_flow.FlowManager[RepairsFlowContext, RepairsFlowResult, str]
):
"""Manage repairs flows."""
@override
async def async_init(
self,
handler: str,
*,
context: RepairsFlowContext | None = None,
data: dict[str, Any] | None = None,
) -> RepairsFlowResult:
"""Override to ensure appropriate context is set in the flow result."""
_context: RepairsFlowContext = context or {}
if "issue_id" not in _context and data is not None and "issue_id" in data:
# fallback for custom integrations
_context |= {"issue_id": data["issue_id"]}
if "issue_id" in _context:
# interim compatibility fallback for custom integrations that may expect
# "issue_id" in user_input of async_step_init
data = {**(data or {}), "issue_id": _context["issue_id"]}
return await super().async_init(handler, context=_context, data=data)
@override
async def async_create_flow(
self,
handler_key: str,
*,
context: data_entry_flow.FlowContext | None = None,
context: RepairsFlowContext | None = None,
data: dict[str, Any] | None = None,
) -> RepairsFlow:
"""Create a flow. platform is a repairs module."""
assert data and "issue_id" in data
issue_id = data["issue_id"]
if context is None or "issue_id" not in context:
raise KeyError("issue_id was not set in context")
issue_id = context["issue_id"]
issue_registry = ir.async_get(self.hass)
issue = issue_registry.async_get_issue(handler_key, issue_id)
@@ -74,16 +94,13 @@ class RepairsFlowManager(
else:
flow = await platform.async_create_fix_flow(self.hass, issue_id, issue.data)
flow.issue_id = issue_id
flow.data = issue.data
return flow
@override
async def async_finish_flow(
self,
flow: data_entry_flow.FlowHandler[
data_entry_flow.FlowContext, RepairsFlowResult, str
],
flow: data_entry_flow.FlowHandler[RepairsFlowContext, RepairsFlowResult, str],
result: RepairsFlowResult,
) -> RepairsFlowResult:
"""Complete a fix flow.
@@ -92,7 +109,7 @@ class RepairsFlowManager(
FlowResultType.CREATE_ENTRY.
"""
if result.get("type") is not data_entry_flow.FlowResultType.ABORT:
ir.async_delete_issue(self.hass, flow.handler, flow.init_data["issue_id"])
ir.async_delete_issue(self.hass, flow.handler, flow.context["issue_id"])
return result
+36 -4
View File
@@ -14,22 +14,54 @@ from homeassistant.core import HomeAssistant, callback
from .const import FlowType
class RepairsFlowContext(data_entry_flow.FlowContext, total=False):
"""Typed flow context for repairs flow."""
issue_id: str
class RepairsFlowResult(
data_entry_flow.FlowResult[data_entry_flow.FlowContext, str], total=False
data_entry_flow.FlowResult[
RepairsFlowContext,
str,
],
total=False,
):
"""Typed result dict for repair flow."""
"""Typed result dict for repairs flow."""
next_flow: tuple[FlowType, str]
result: ConfigEntry | None
class RepairsFlow(
data_entry_flow.FlowHandler[data_entry_flow.FlowContext, RepairsFlowResult, str]
data_entry_flow.FlowHandler[
RepairsFlowContext,
RepairsFlowResult,
str,
]
):
"""Handle a flow for fixing an issue."""
issue_id: str
data: dict[str, str | int | float | None] | None
_issue_id: str
@property
def issue_id(self) -> str:
"""Return the flow's issue_id."""
if "issue_id" in self.context:
return self.context["issue_id"]
# Avoid breaking changes in legacy custom integrations that may access
# this property prior to the flow manager applying the context in async_create_flow.
return self._issue_id
@issue_id.setter
def issue_id(self, issue_id: str) -> None:
"""Allow legacy implementations to set issue_id.
Setter is retained to avoid breaking changes in custom integrations that may set issue_id in a RepairFlow
prior to the flow manager applying the context.
"""
self._issue_id = issue_id
@override
@callback
@@ -144,7 +144,7 @@ class RepairsFlowIndexView(FlowManagerIndexView[RepairsFlowManager, RepairsFlowR
try:
result = await self._flow_mgr.async_init(
data["handler"],
data={"issue_id": data["issue_id"]},
context={"issue_id": data["issue_id"]},
)
except data_entry_flow.UnknownFlow as ex:
return self.json_message(
@@ -0,0 +1,97 @@
"""Tests for repairs issue_handler.py."""
import pytest
from homeassistant.components.repairs import (
DOMAIN,
RepairsFlow,
RepairsFlowResult,
repairs_flow_manager,
)
from homeassistant.core import HomeAssistant
import homeassistant.helpers.issue_registry as ir
from tests.common import AsyncMock, Mock, async_setup_component, mock_platform
@pytest.fixture(autouse=True)
async def mock_repairs_integration(hass: HomeAssistant) -> None:
"""Mock a repairs integration."""
hass.config.components.add("fake_integration")
async def async_create_fix_flow(
hass: HomeAssistant,
issue_id: str,
data: dict[str, str | int | float | None] | None,
) -> RepairsFlow:
return MockFixFlowContext()
mock_platform(
hass,
"fake_integration.repairs",
Mock(async_create_fix_flow=AsyncMock(wraps=async_create_fix_flow)),
)
class MockFixFlowContext(RepairsFlow):
"""Mock for context tests."""
def __init__(self) -> None:
"""Initialize a MockFlowFixContext."""
# Test issue_id setter
self.issue_id = "fake_issue"
assert self.issue_id == "fake_issue"
async def async_step_init(self, user_input: dict | None) -> RepairsFlowResult:
"""Initial step of a repairs flow."""
assert user_input and user_input["issue_id"] == self.issue_id
return self.async_show_form()
@pytest.mark.parametrize(
("ignore_translations_for_mock_domains"),
[
["fake_integration"],
],
)
async def test_flow_fix_via_data(hass: HomeAssistant) -> None:
"""Test that a repairs flow's issue_id can be set via data."""
assert await async_setup_component(hass, DOMAIN, {})
ir.async_create_issue(
hass,
issue_id="context_issue",
domain="fake_integration",
is_fixable=True,
severity="error",
translation_key="fake_key",
)
assert (repairs := repairs_flow_manager(hass))
result = await repairs.async_init(
"fake_integration", data={"issue_id": "context_issue"}
)
assert result["type"] == "form"
result = repairs.async_get(result["flow_id"])
assert result["context"] == {"issue_id": "context_issue"}
@pytest.mark.parametrize(
("ignore_translations_for_mock_domains"),
[
["fake_integration"],
],
)
async def test_flow_fix_missing_context(hass: HomeAssistant) -> None:
"""Test that KeyError is thrown when context and data is missing."""
assert await async_setup_component(hass, DOMAIN, {})
assert (repairs := repairs_flow_manager(hass))
with pytest.raises(KeyError) as exc:
await repairs.async_init("fake_integration")
assert "issue_id was not set in context" in str(exc.value)
+38 -2
View File
@@ -22,6 +22,8 @@ from homeassistant.config_entries import (
from homeassistant.core import HomeAssistant, callback
import homeassistant.helpers.issue_registry as ir
from .test_issue_handler import MockFixFlowContext
from tests.common import (
AsyncMock,
Mock,
@@ -44,6 +46,8 @@ async def mock_repairs_integration(hass: HomeAssistant) -> None:
issue_id: str,
data: dict[str, str | int | float | None] | None,
) -> RepairsFlow:
if issue_id == "context_issue":
return MockFixFlowContext()
return MockFixFlowNextFlow()
mock_platform(
@@ -117,7 +121,6 @@ class MockFixFlowNextFlow(RepairsFlow):
return self.async_create_entry(
next_flow=(FlowType.OPTIONS_FLOW, next_flow["flow_id"]), data={}
)
# self.issue_id == "subentry_config_issue"
assert len(mock_entry.subentries) == 1
next_flow = await self.hass.config_entries.subentries.async_init(
(mock_entry.entry_id, "fake_subentry"),
@@ -169,9 +172,42 @@ async def test_fix_issue_next_flow(hass: HomeAssistant, flow_type: FlowType) ->
assert (repairs := repairs_flow_manager(hass))
flow = await repairs.async_init("fake_integration", data={"issue_id": flow_type})
flow = await repairs.async_init(
"fake_integration", context={"issue_id": str(flow_type)}
)
next_flow_type, _ = flow["next_flow"]
assert next_flow_type is flow_type
assert mock_entry == flow["result"]
@pytest.mark.parametrize(
("ignore_translations_for_mock_domains"),
[
["fake_integration"],
],
)
async def test_issue_id_setter_getter(hass: HomeAssistant) -> None:
"""Test RepairFlow issue_id getter/setter with switch to context."""
assert await async_setup_component(hass, DOMAIN, {})
ir.async_create_issue(
hass,
issue_id="context_issue",
domain="fake_integration",
is_fixable=True,
severity="error",
translation_key="fake_key",
)
assert (repairs := repairs_flow_manager(hass))
result = await repairs.async_init(
"fake_integration", context={"issue_id": "context_issue"}
)
assert result["type"] == "form"
result = repairs.async_get(result["flow_id"])
assert result["context"] == {"issue_id": "context_issue"}