diff --git a/homeassistant/components/mqtt/__init__.py b/homeassistant/components/mqtt/__init__.py index acec4613d4f9..fb161f17a6f1 100644 --- a/homeassistant/components/mqtt/__init__.py +++ b/homeassistant/components/mqtt/__init__.py @@ -2,7 +2,6 @@ import asyncio from collections.abc import Callable -from datetime import datetime import logging from typing import Any, cast @@ -18,24 +17,16 @@ from homeassistant.const import ( CONF_PROTOCOL, SERVICE_RELOAD, ) -from homeassistant.core import HomeAssistant, ServiceCall, callback -from homeassistant.exceptions import ( - ConfigValidationError, - ServiceValidationError, - Unauthorized, -) +from homeassistant.core import HomeAssistant, callback +from homeassistant.exceptions import Unauthorized from homeassistant.helpers import ( config_validation as cv, entity_registry as er, - event as ev, issue_registry as ir, ) from homeassistant.helpers.device_registry import AnyDeviceEntry from homeassistant.helpers.dispatcher import async_dispatcher_connect -from homeassistant.helpers.entity_platform import async_get_platforms from homeassistant.helpers.issue_registry import IssueSeverity, async_create_issue -from homeassistant.helpers.reload import async_integration_yaml_config -from homeassistant.helpers.service import async_register_admin_service from homeassistant.helpers.typing import ConfigType from homeassistant.loader import async_get_loaded_integration from homeassistant.setup import SetupPhases, async_pause_setup @@ -56,7 +47,6 @@ from .client import ( from .config import MQTT_BASE_SCHEMA, MQTT_RO_SCHEMA, MQTT_RW_SCHEMA from .config_integration import CONFIG_SCHEMA_BASE from .const import ( - ATTR_MESSAGE_EXPIRY_INTERVAL, ATTR_PAYLOAD, ATTR_QOS, ATTR_RETAIN, @@ -90,6 +80,7 @@ from .const import ( MQTT_CONNECTION_STATE, PROTOCOL_5, PROTOCOL_311, + SERVICE_PUBLISH, TEMPLATE_ERRORS, Platform, ) @@ -104,6 +95,7 @@ from .models import ( ReceiveMessage, convert_outgoing_mqtt_payload, ) +from .services import async_setup_services from .subscription import ( EntitySubscription, async_prepare_subscribe_topics, @@ -160,6 +152,7 @@ __all__ = [ "MQTT_CONNECTION_STATE", "MQTT_RO_SCHEMA", "MQTT_RW_SCHEMA", + "SERVICE_PUBLISH", "SERVICE_RELOAD", "TEMPLATE_ERRORS", "EntitySubscription", @@ -199,11 +192,6 @@ __all__ = [ _LOGGER = logging.getLogger(__name__) -SERVICE_PUBLISH = "publish" -SERVICE_DUMP = "dump" - -ATTR_EVALUATE_PAYLOAD = "evaluate_payload" - MAX_RECONNECT_WAIT = 300 # seconds CONNECTION_SUCCESS = "connection_success" @@ -242,19 +230,6 @@ CONFIG_SCHEMA = probatio.Schema( extra=probatio.ALLOW_EXTRA, ) -# Publish action call validation schema -MQTT_PUBLISH_SCHEMA = probatio.Schema( - { - probatio.Required(ATTR_TOPIC): valid_publish_topic, - probatio.Required(ATTR_PAYLOAD, default=None): probatio.Any(cv.string, None), - probatio.Optional(ATTR_EVALUATE_PAYLOAD): cv.boolean, - probatio.Optional(ATTR_QOS, default=DEFAULT_QOS): valid_qos_schema, - probatio.Optional(ATTR_RETAIN, default=DEFAULT_RETAIN): cv.boolean, - probatio.Optional(ATTR_MESSAGE_EXPIRY_INTERVAL): cv.positive_time_period_dict, - }, - required=True, -) - async def _async_config_entry_updated(hass: HomeAssistant, entry: ConfigEntry) -> None: """Handle signals of config entry being updated. @@ -298,132 +273,7 @@ async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool: websocket_api.async_register_command(hass, websocket_subscribe) websocket_api.async_register_command(hass, websocket_mqtt_info) - async def async_publish_service(call: ServiceCall) -> None: - """Handle MQTT publish service calls.""" - msg_topic: str = call.data[ATTR_TOPIC] - - if not mqtt_config_entry_enabled(hass): - raise ServiceValidationError( - translation_key="mqtt_not_setup_cannot_publish", - translation_domain=DOMAIN, - translation_placeholders={"topic": msg_topic}, - ) - - mqtt_data = hass.data[DATA_MQTT] - payload: PublishPayloadType = call.data[ATTR_PAYLOAD] - evaluate_payload: bool = call.data.get(ATTR_EVALUATE_PAYLOAD, False) - qos: int = call.data[ATTR_QOS] - retain: bool = call.data[ATTR_RETAIN] - message_expiry_interval: int | None = ( - int(call.data[ATTR_MESSAGE_EXPIRY_INTERVAL].total_seconds()) - if ATTR_MESSAGE_EXPIRY_INTERVAL in call.data - else None - ) - - if evaluate_payload: - # Convert quoted binary literal to raw data - payload = convert_outgoing_mqtt_payload(payload) - - await mqtt_data.client.async_publish( - msg_topic, - payload, - qos, - retain, - message_expiry_interval=message_expiry_interval, - ) - - async_register_admin_service( - hass, DOMAIN, SERVICE_PUBLISH, async_publish_service, MQTT_PUBLISH_SCHEMA - ) - - async def async_dump_service(call: ServiceCall) -> None: - """Handle MQTT dump service calls.""" - messages: list[tuple[str, str]] = [] - - @callback - def collect_msg(msg: ReceiveMessage) -> None: - messages.append((msg.topic, str(msg.payload).replace("\n", ""))) - - unsub = async_subscribe_internal(hass, call.data["topic"], collect_msg) - - def write_dump() -> None: - with open(hass.config.path("mqtt_dump.txt"), "w", encoding="utf8") as fp: - fp.writelines([",".join(msg) + "\n" for msg in messages]) - - async def finish_dump(_: datetime) -> None: - """Write dump to file.""" - unsub() - await hass.async_add_executor_job(write_dump) - - ev.async_call_later(hass, call.data["duration"], finish_dump) - - async_register_admin_service( - hass, - DOMAIN, - SERVICE_DUMP, - async_dump_service, - schema=probatio.Schema( - { - probatio.Required("topic"): valid_subscribe_topic, - probatio.Optional("duration", default=5): int, - } - ), - ) - - async def _reload_config(call: ServiceCall) -> None: - """Reload the platforms.""" - if not mqtt_config_entry_enabled(hass): - _LOGGER.debug( - "Skipped reloading MQTT integration, " - "the MQTT config entry is not enabled" - ) - return - entry: ConfigEntry = next(iter(hass.config_entries.async_entries(DOMAIN))) - mqtt_data = hass.data[DATA_MQTT] - - # Fetch updated manually configured items and validate - try: - config_yaml = await async_integration_yaml_config( - hass, DOMAIN, raise_on_failure=True - ) - except ConfigValidationError as ex: - raise ServiceValidationError( - translation_domain=ex.translation_domain, - translation_key=ex.translation_key, - translation_placeholders=ex.translation_placeholders, - ) from ex - - new_config: list[ConfigType] = config_yaml.get(DOMAIN, []) - platforms_used = platforms_from_config(new_config) - new_platforms = platforms_used - mqtt_data.platforms_loaded - await async_forward_entry_setup_and_setup_discovery(hass, entry, new_platforms) - # Check the schema before continuing reload - await async_check_config_schema(hass, config_yaml) - - # Remove repair issues - async_remove_mqtt_issues(hass, mqtt_data) - - mqtt_data.config = new_config - - # Reload the modern yaml platforms - mqtt_platforms = async_get_platforms(hass, DOMAIN) - tasks = [ - create_eager_task(entity.async_remove()) - for mqtt_platform in mqtt_platforms - for entity in list(mqtt_platform.entities.values()) - if getattr(entity, "_discovery_data", None) is None - and mqtt_platform.config_entry - and mqtt_platform.domain in ENTITY_PLATFORMS - ] - await asyncio.gather(*tasks) - - for component in mqtt_data.reload_handlers.values(): - component() - - # Fire event - hass.bus.async_fire(f"event_{DOMAIN}_reloaded", context=call.context) - - async_register_admin_service(hass, DOMAIN, SERVICE_RELOAD, _reload_config) + async_setup_services(hass) return True diff --git a/homeassistant/components/mqtt/const.py b/homeassistant/components/mqtt/const.py index f498a57769bf..c07335496b67 100644 --- a/homeassistant/components/mqtt/const.py +++ b/homeassistant/components/mqtt/const.py @@ -11,6 +11,7 @@ from homeassistant.exceptions import TemplateError ATTR_DISCOVERY_HASH = "discovery_hash" ATTR_DISCOVERY_PAYLOAD = "discovery_payload" ATTR_DISCOVERY_TOPIC = "discovery_topic" +ATTR_EVALUATE_PAYLOAD = "evaluate_payload" ATTR_MESSAGE_EXPIRY_INTERVAL = "message_expiry_interval" ATTR_PAYLOAD = "payload" ATTR_QOS = "qos" @@ -382,6 +383,9 @@ MQTT_PROCESSED_SUBSCRIPTIONS = "mqtt_processed_subscriptions" PAYLOAD_EMPTY_JSON = "{}" PAYLOAD_NONE = "None" +SERVICE_DUMP = "dump" +SERVICE_PUBLISH = "publish" + CONFIG_ENTRY_VERSION = 2 CONFIG_ENTRY_MINOR_VERSION = 1 diff --git a/homeassistant/components/mqtt/services.py b/homeassistant/components/mqtt/services.py new file mode 100644 index 000000000000..dfc30ee18587 --- /dev/null +++ b/homeassistant/components/mqtt/services.py @@ -0,0 +1,196 @@ +"""Support for MQTT actions.""" + +import asyncio +from datetime import datetime + +import probatio + +from homeassistant.config_entries import ConfigEntry +from homeassistant.const import SERVICE_RELOAD +from homeassistant.core import HomeAssistant, ServiceCall, callback +from homeassistant.exceptions import ConfigValidationError, ServiceValidationError +from homeassistant.helpers import config_validation as cv, event as ev +from homeassistant.helpers.entity_platform import async_get_platforms +from homeassistant.helpers.reload import async_integration_yaml_config +from homeassistant.helpers.service import async_register_admin_service +from homeassistant.helpers.typing import ConfigType +from homeassistant.util.async_ import create_eager_task + +from .client import async_subscribe_internal +from .const import ( + ATTR_EVALUATE_PAYLOAD, + ATTR_MESSAGE_EXPIRY_INTERVAL, + ATTR_PAYLOAD, + ATTR_QOS, + ATTR_RETAIN, + ATTR_TOPIC, + DEFAULT_QOS, + DEFAULT_RETAIN, + DOMAIN, + ENTITY_PLATFORMS, + LOGGER, + SERVICE_DUMP, + SERVICE_PUBLISH, +) +from .models import ( + DATA_MQTT, + PublishPayloadType, + ReceiveMessage, + convert_outgoing_mqtt_payload, +) +from .util import ( + async_check_config_schema, + async_forward_entry_setup_and_setup_discovery, + async_remove_mqtt_issues, + mqtt_config_entry_enabled, + platforms_from_config, + valid_publish_topic, + valid_qos_schema, + valid_subscribe_topic, +) + +# Publish action call validation schema +MQTT_PUBLISH_SCHEMA = probatio.Schema( + { + probatio.Required(ATTR_TOPIC): valid_publish_topic, + probatio.Required(ATTR_PAYLOAD, default=None): probatio.Any(cv.string, None), + probatio.Optional(ATTR_EVALUATE_PAYLOAD): cv.boolean, + probatio.Optional(ATTR_QOS, default=DEFAULT_QOS): valid_qos_schema, + probatio.Optional(ATTR_RETAIN, default=DEFAULT_RETAIN): cv.boolean, + probatio.Optional(ATTR_MESSAGE_EXPIRY_INTERVAL): cv.positive_time_period_dict, + }, + required=True, +) + +MQTT_DUMP_SCHEMA = probatio.Schema( + { + probatio.Required("topic"): valid_subscribe_topic, + probatio.Optional("duration", default=5): int, + } +) + + +async def _async_publish_service(call: ServiceCall) -> None: + """Handle MQTT publish service calls.""" + hass = call.hass + msg_topic: str = call.data[ATTR_TOPIC] + + if not mqtt_config_entry_enabled(hass): + raise ServiceValidationError( + translation_key="mqtt_not_setup_cannot_publish", + translation_domain=DOMAIN, + translation_placeholders={"topic": msg_topic}, + ) + + mqtt_data = hass.data[DATA_MQTT] + payload: PublishPayloadType = call.data[ATTR_PAYLOAD] + evaluate_payload: bool = call.data.get(ATTR_EVALUATE_PAYLOAD, False) + qos: int = call.data[ATTR_QOS] + retain: bool = call.data[ATTR_RETAIN] + message_expiry_interval: int | None = ( + int(call.data[ATTR_MESSAGE_EXPIRY_INTERVAL].total_seconds()) + if ATTR_MESSAGE_EXPIRY_INTERVAL in call.data + else None + ) + + if evaluate_payload: + # Convert quoted binary literal to raw data + payload = convert_outgoing_mqtt_payload(payload) + + await mqtt_data.client.async_publish( + msg_topic, + payload, + qos, + retain, + message_expiry_interval=message_expiry_interval, + ) + + +async def _async_dump_service(call: ServiceCall) -> None: + """Handle MQTT dump service calls.""" + hass = call.hass + messages: list[tuple[str, str]] = [] + + @callback + def collect_msg(msg: ReceiveMessage) -> None: + messages.append((msg.topic, str(msg.payload).replace("\n", ""))) + + unsub = async_subscribe_internal(hass, call.data["topic"], collect_msg) + + def write_dump() -> None: + with open(hass.config.path("mqtt_dump.txt"), "w", encoding="utf8") as fp: + fp.writelines([",".join(msg) + "\n" for msg in messages]) + + async def finish_dump(_: datetime) -> None: + """Write dump to file.""" + unsub() + await hass.async_add_executor_job(write_dump) + + ev.async_call_later(hass, call.data["duration"], finish_dump) + + +async def _async_reload_config(call: ServiceCall) -> None: + """Reload the platforms.""" + hass = call.hass + if not mqtt_config_entry_enabled(hass): + LOGGER.debug( + "Skipped reloading MQTT integration, the MQTT config entry is not enabled" + ) + return + entry: ConfigEntry = next(iter(hass.config_entries.async_entries(DOMAIN))) + mqtt_data = hass.data[DATA_MQTT] + + # Fetch updated manually configured items and validate + try: + config_yaml = await async_integration_yaml_config( + hass, DOMAIN, raise_on_failure=True + ) + except ConfigValidationError as ex: + raise ServiceValidationError( + translation_domain=ex.translation_domain, + translation_key=ex.translation_key, + translation_placeholders=ex.translation_placeholders, + ) from ex + + new_config: list[ConfigType] = config_yaml.get(DOMAIN, []) + platforms_used = platforms_from_config(new_config) + new_platforms = platforms_used - mqtt_data.platforms_loaded + await async_forward_entry_setup_and_setup_discovery(hass, entry, new_platforms) + # Check the schema before continuing reload + await async_check_config_schema(hass, config_yaml) + + # Remove repair issues + async_remove_mqtt_issues(hass, mqtt_data) + + mqtt_data.config = new_config + + # Reload the modern yaml platforms + mqtt_platforms = async_get_platforms(hass, DOMAIN) + tasks = [ + create_eager_task(entity.async_remove()) + for mqtt_platform in mqtt_platforms + for entity in list(mqtt_platform.entities.values()) + if getattr(entity, "_discovery_data", None) is None + and mqtt_platform.config_entry + and mqtt_platform.domain in ENTITY_PLATFORMS + ] + await asyncio.gather(*tasks) + + for component in mqtt_data.reload_handlers.values(): + component() + + # Fire event + hass.bus.async_fire(f"event_{DOMAIN}_reloaded", context=call.context) + + +@callback +def async_setup_services(hass: HomeAssistant) -> None: + """Set up the actions for the MQTT component.""" + + async_register_admin_service( + hass, DOMAIN, SERVICE_PUBLISH, _async_publish_service, MQTT_PUBLISH_SCHEMA + ) + async_register_admin_service( + hass, DOMAIN, SERVICE_DUMP, _async_dump_service, MQTT_DUMP_SCHEMA + ) + async_register_admin_service(hass, DOMAIN, SERVICE_RELOAD, _async_reload_config) diff --git a/tests/components/mqtt/test_init.py b/tests/components/mqtt/test_init.py index 338c63869379..add36f342475 100644 --- a/tests/components/mqtt/test_init.py +++ b/tests/components/mqtt/test_init.py @@ -18,7 +18,12 @@ import pytest from homeassistant import core as ha from homeassistant.components import mqtt from homeassistant.components.mqtt import debug_info -from homeassistant.components.mqtt.const import DOMAIN +from homeassistant.components.mqtt.const import ( + ATTR_EVALUATE_PAYLOAD, + ATTR_MESSAGE_EXPIRY_INTERVAL, + DOMAIN, + SERVICE_DUMP, +) from homeassistant.components.mqtt.models import ( MessageCallbackType, MqttCommandTemplateException, @@ -365,7 +370,7 @@ async def test_mqtt_publish_action_call_with_raw_data( { mqtt.ATTR_TOPIC: "test/topic", mqtt.ATTR_PAYLOAD: attr_payload, - mqtt.ATTR_EVALUATE_PAYLOAD: evaluate_payload, + ATTR_EVALUATE_PAYLOAD: evaluate_payload, }, blocking=True, ) @@ -392,7 +397,7 @@ async def test_mqtt_publish_action_call_with_raw_data( { mqtt.ATTR_TOPIC: "test/topic", mqtt.ATTR_PAYLOAD: attr_payload, - mqtt.ATTR_EVALUATE_PAYLOAD: evaluate_payload, + ATTR_EVALUATE_PAYLOAD: evaluate_payload, }, blocking=True, ) @@ -496,7 +501,7 @@ async def test_publish_action_with_message_expiry_interval( mqtt.ATTR_PAYLOAD: "bla", mqtt.ATTR_QOS: 2, mqtt.ATTR_RETAIN: False, - mqtt.ATTR_MESSAGE_EXPIRY_INTERVAL: interval_data, + ATTR_MESSAGE_EXPIRY_INTERVAL: interval_data, }, blocking=True, ) @@ -1112,7 +1117,7 @@ async def test_dump_service( async_fire_mqtt_message(hass, "bla/1", "test1") async_fire_mqtt_message(hass, "bla/2", "test2") - with patch("homeassistant.components.mqtt.open", mopen): + with patch("homeassistant.components.mqtt.services.open", mopen): async_fire_time_changed(hass, utcnow() + timedelta(seconds=3)) await hass.async_block_till_done() @@ -1127,7 +1132,7 @@ ADMIN_SERVICE_CALLS = [ {mqtt.ATTR_TOPIC: "test/topic", mqtt.ATTR_PAYLOAD: "payload"}, id="publish", ), - pytest.param(mqtt.SERVICE_DUMP, {"topic": "bla/#", "duration": 3}, id="dump"), + pytest.param(SERVICE_DUMP, {"topic": "bla/#", "duration": 3}, id="dump"), ]