Move MQTT actions to services module (#184326)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
epenet
2026-10-05 16:59:17 +02:00
committed by GitHub
co-authored by Claude Opus 5
parent ef31c2d9ca
commit 859b73d71f
4 changed files with 217 additions and 162 deletions
+6 -156
View File
@@ -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
+4
View File
@@ -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
+196
View File
@@ -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)
+11 -6
View File
@@ -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"),
]