Only import MQTT in MySensors once it is set up (#183290)

Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
Franck Nijhof
2026-09-28 13:15:18 +02:00
committed by GitHub
co-authored by Claude
parent 0ff182c6fa
commit f3628f2420
4 changed files with 69 additions and 12 deletions
@@ -10,11 +10,6 @@ from awesomeversion import (
)
import probatio
from homeassistant.components.mqtt import (
DOMAIN as MQTT_DOMAIN,
valid_publish_topic,
valid_subscribe_topic,
)
from homeassistant.config_entries import ConfigEntry, ConfigFlow, ConfigFlowResult
from homeassistant.const import CONF_DEVICE
from homeassistant.core import callback
@@ -34,6 +29,7 @@ from .const import (
CONF_TOPIC_OUT_PREFIX,
CONF_VERSION,
DOMAIN,
MQTT_DOMAIN,
ConfGatewayType,
)
from .gateway import MQTT_COMPONENT, is_serial_port, is_socket_address, try_connect
@@ -211,6 +207,11 @@ class MySensorsConfigFlowHandler(ConfigFlow, domain=DOMAIN):
if MQTT_DOMAIN not in self.hass.config.components:
return self.async_abort(reason="mqtt_required")
from homeassistant.components.mqtt import ( # noqa: PLC0415
valid_publish_topic,
valid_subscribe_topic,
)
gw_type = self._gw_type = CONF_GATEWAY_TYPE_MQTT
errors: dict[str, str] = {}
@@ -5,6 +5,9 @@ from typing import Final, Literal, TypedDict
from homeassistant.const import Platform
# MQTT is only imported once it is set up, as it is heavy to load
MQTT_DOMAIN: Final = "mqtt"
ATTR_DEVICES: Final = "devices"
ATTR_GATEWAY_ID: Final = "gateway_id"
ATTR_NODE_ID: Final = "node_id"
@@ -11,12 +11,6 @@ from typing import Any
from mysensors import BaseAsyncGateway, Message, Sensor, get_const, mysensors
import probatio
from homeassistant.components.mqtt import (
DOMAIN as MQTT_DOMAIN,
ReceiveMessage as MQTTReceiveMessage,
async_publish,
async_subscribe,
)
from homeassistant.const import CONF_DEVICE, EVENT_HOMEASSISTANT_STOP
from homeassistant.core import Event, HomeAssistant, callback
from homeassistant.helpers import config_validation as cv
@@ -35,6 +29,7 @@ from .const import (
CONF_TOPIC_IN_PREFIX,
CONF_TOPIC_OUT_PREFIX,
CONF_VERSION,
MQTT_DOMAIN,
ConfGatewayType,
)
from .handler import HANDLERS
@@ -176,6 +171,12 @@ async def _get_gateway(
if MQTT_DOMAIN not in hass.config.components:
return None
from homeassistant.components.mqtt import ( # noqa: PLC0415
ReceiveMessage as MQTTReceiveMessage,
async_publish,
async_subscribe,
)
def pub_callback(topic: str, payload: str, qos: int, retain: bool) -> None:
"""Call MQTT publish function."""
hass.async_create_task(async_publish(hass, topic, payload, qos, retain))
+53 -1
View File
@@ -1,13 +1,26 @@
"""Test function in gateway.py."""
from unittest.mock import patch
from unittest.mock import AsyncMock, MagicMock, patch
import probatio
import pytest
from homeassistant.components.mysensors.const import (
CONF_GATEWAY_TYPE,
CONF_GATEWAY_TYPE_MQTT,
CONF_RETAIN,
CONF_TOPIC_IN_PREFIX,
CONF_TOPIC_OUT_PREFIX,
CONF_VERSION,
DOMAIN,
)
from homeassistant.components.mysensors.gateway import is_serial_port
from homeassistant.const import CONF_DEVICE
from homeassistant.core import HomeAssistant
from tests.common import MockConfigEntry, async_fire_mqtt_message
from tests.typing import MqttMockHAClient
@pytest.mark.parametrize(
("port", "expect_valid"),
@@ -31,3 +44,42 @@ def test_is_serial_port_windows(
assert not expect_valid
else:
assert expect_valid
async def test_mqtt_gateway(hass: HomeAssistant, mqtt_mock: MqttMockHAClient) -> None:
"""Test the MQTT gateway subscribes and publishes through MQTT."""
entry = MockConfigEntry(
domain=DOMAIN,
data={
CONF_GATEWAY_TYPE: CONF_GATEWAY_TYPE_MQTT,
CONF_DEVICE: "mqtt",
CONF_VERSION: "2.3",
CONF_TOPIC_IN_PREFIX: "in",
CONF_TOPIC_OUT_PREFIX: "out",
CONF_RETAIN: False,
},
)
entry.add_to_hass(hass)
with (
patch("mysensors.task.OTAFirmware", autospec=True),
patch("mysensors.task.load_fw", autospec=True),
patch("mysensors.task.Persistence", autospec=True) as persistence_class,
):
persistence = persistence_class.return_value
persistence.schedule_save_sensors = AsyncMock()
persistence.safe_load_sensors = MagicMock()
persistence.save_sensors = MagicMock()
assert await hass.config_entries.async_setup(entry.entry_id)
await hass.async_block_till_done()
subscribed_topics = [
call.args[0] for call in mqtt_mock.async_subscribe.call_args_list
]
assert "in/+/+/3/+/+" in subscribed_topics
# A time request is answered by publishing the current time
async_fire_mqtt_message(hass, "in/1/255/3/0/1", "")
await hass.async_block_till_done()
published_topics = [call.args[0] for call in mqtt_mock.async_publish.call_args_list]
assert "out/1/255/3/0/1" in published_topics