From 78083284b97145b191b4f3f09bd676fc699ad0d1 Mon Sep 17 00:00:00 2001 From: epenet <6771947+epenet@users.noreply.github.com> Date: Thu, 17 Sep 2026 18:16:01 +0200 Subject: [PATCH] Move assist_satellite service registration to services module (#182463) --- .../components/assist_satellite/__init__.py | 188 +--------------- .../components/assist_satellite/services.py | 200 ++++++++++++++++++ 2 files changed, 203 insertions(+), 185 deletions(-) create mode 100644 homeassistant/components/assist_satellite/services.py diff --git a/homeassistant/components/assist_satellite/__init__.py b/homeassistant/components/assist_satellite/__init__.py index c353f9ad1926..2bb1ba7c9e04 100644 --- a/homeassistant/components/assist_satellite/__init__.py +++ b/homeassistant/components/assist_satellite/__init__.py @@ -1,27 +1,11 @@ """Base class for assist satellite entities.""" -from dataclasses import asdict import logging from pathlib import Path -import re -from typing import Any -from hassil.parse_expression import parse_sentence -from hassil.parser import ParseError -from hassil.util import ( - PUNCTUATION_END, - PUNCTUATION_END_WORD, - PUNCTUATION_START, - PUNCTUATION_START_WORD, -) -import probatio - -from homeassistant.auth.permissions.const import CAT_ENTITIES, POLICY_CONTROL from homeassistant.components.http import StaticPathConfig from homeassistant.config_entries import ConfigEntry -from homeassistant.const import ATTR_ENTITY_ID -from homeassistant.core import HomeAssistant, ServiceCall, SupportsResponse -from homeassistant.exceptions import HomeAssistantError, Unauthorized, UnknownUser +from homeassistant.core import HomeAssistant from homeassistant.helpers import config_validation as cv from homeassistant.helpers.entity_component import EntityComponent from homeassistant.helpers.typing import ConfigType @@ -44,6 +28,7 @@ from .entity import ( AssistSatelliteWakeWord, ) from .errors import SatelliteBusyError +from .services import async_setup_services from .websocket_api import async_register_websocket_api __all__ = [ @@ -69,115 +54,8 @@ async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool: ) await component.async_setup(config) - component.async_register_entity_service( - "announce", - probatio.All( - cv.make_entity_service_schema( - { - probatio.Optional("message"): str, - probatio.Optional("media_id"): _media_id_validator, - probatio.Optional("preannounce", default=True): bool, - probatio.Optional("preannounce_media_id"): _media_id_validator, - } - ), - cv.has_at_least_one_key("message", "media_id"), - ), - "async_internal_announce", - [AssistSatelliteEntityFeature.ANNOUNCE], - ) + async_setup_services(hass) - component.async_register_entity_service( - "start_conversation", - probatio.All( - cv.make_entity_service_schema( - { - probatio.Optional("start_message"): str, - probatio.Optional("start_media_id"): _media_id_validator, - probatio.Optional("preannounce", default=True): bool, - probatio.Optional("preannounce_media_id"): _media_id_validator, - probatio.Optional("extra_system_prompt"): str, - } - ), - cv.has_at_least_one_key("start_message", "start_media_id"), - ), - "async_internal_start_conversation", - [AssistSatelliteEntityFeature.START_CONVERSATION], - ) - - async def handle_ask_question(call: ServiceCall) -> dict[str, Any]: - """Handle a Show View service call.""" - satellite_entity_id: str = call.data[ATTR_ENTITY_ID] - if call.context.user_id: - user = await hass.auth.async_get_user(call.context.user_id) - if user is None: - raise UnknownUser( - context=call.context, - permission=POLICY_CONTROL, - user_id=call.context.user_id, - ) - if not user.permissions.check_entity(satellite_entity_id, POLICY_CONTROL): - raise Unauthorized( - context=call.context, - permission=POLICY_CONTROL, - user_id=call.context.user_id, - perm_category=CAT_ENTITIES, - ) - - satellite_entity: AssistSatelliteEntity | None = component.get_entity( - satellite_entity_id - ) - if satellite_entity is None: - raise HomeAssistantError( - f"Invalid Assist satellite entity id: {satellite_entity_id}" - ) - - satellite_entity.async_set_context(call.context) - - ask_question_args = { - "question": call.data.get("question"), - "question_media_id": call.data.get("question_media_id"), - "preannounce": call.data.get("preannounce", True), - "answers": call.data.get("answers"), - } - - if preannounce_media_id := call.data.get("preannounce_media_id"): - ask_question_args["preannounce_media_id"] = preannounce_media_id - - answer = await satellite_entity.async_internal_ask_question(**ask_question_args) - - if answer is None: - raise HomeAssistantError("No answer from satellite") - - return asdict(answer) - - hass.services.async_register( - domain=DOMAIN, - service="ask_question", - service_func=handle_ask_question, - schema=probatio.All( - { - probatio.Required(ATTR_ENTITY_ID): cv.entity_domain(DOMAIN), - probatio.Optional("question"): str, - probatio.Optional("question_media_id"): _media_id_validator, - probatio.Optional("preannounce", default=True): bool, - probatio.Optional("preannounce_media_id"): _media_id_validator, - probatio.Optional("answers"): [ - { - probatio.Required("id"): str, - probatio.Required("sentences"): probatio.All( - cv.ensure_list, - [cv.string], - has_one_non_empty_item, - has_no_punctuation, - is_valid_sentence, - ), - } - ], - }, - cv.has_at_least_one_key("question", "question_media_id"), - ), - supports_response=SupportsResponse.ONLY, - ) hass.data[CONNECTION_TEST_DATA] = {} async_register_websocket_api(hass) hass.http.register_view(ConnectionTestView()) @@ -202,63 +80,3 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool: async def async_unload_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool: """Unload a config entry.""" return await hass.data[DATA_COMPONENT].async_unload_entry(entry) - - -def has_no_punctuation(value: list[str]) -> list[str]: - """Validate result does not contain punctuation.""" - for sentence in value: - # Exclude {list_references} which may contain punctuation characters. - sentence = _remove_list_references(sentence) - if ( - PUNCTUATION_START.search(sentence) - or PUNCTUATION_END.search(sentence) - or PUNCTUATION_START_WORD.search(sentence) - or PUNCTUATION_END_WORD.search(sentence) - ): - raise probatio.Invalid("sentence should not contain punctuation") - - return value - - -def _remove_list_references(sentence: str) -> str: - """Remove {list_references} from a sentence for linting.""" - return re.sub(r"(? list[str]: - """Validate result can be parsed by hassil.""" - for sentence in value: - try: - parse_sentence(sentence) - except ParseError as err: - raise probatio.Invalid(f"invalid sentence: {err}") from err - return value - - -def has_one_non_empty_item(value: list[str]) -> list[str]: - """Validate result has at least one item.""" - if len(value) < 1: - raise probatio.Invalid("at least one sentence is required") - - for sentence in value: - if not sentence: - raise probatio.Invalid("sentences cannot be empty") - - return value - - -# Validator for media_id fields that accepts both string and media selector format -_media_id_validator = probatio.Any( - cv.string, # Plain string format - probatio.All( - probatio.Schema( - { - probatio.Required("media_content_id"): cv.string, - probatio.Required("media_content_type"): cv.string, - probatio.Remove("metadata"): dict, # Ignore metadata if present - } - ), - # Extract media_content_id from media selector format - lambda x: x["media_content_id"], - ), -) diff --git a/homeassistant/components/assist_satellite/services.py b/homeassistant/components/assist_satellite/services.py new file mode 100644 index 000000000000..e38a21810d85 --- /dev/null +++ b/homeassistant/components/assist_satellite/services.py @@ -0,0 +1,200 @@ +"""Services for the Assist satellite integration.""" + +from dataclasses import asdict +import re +from typing import Any + +from hassil.parse_expression import parse_sentence +from hassil.parser import ParseError +from hassil.util import ( + PUNCTUATION_END, + PUNCTUATION_END_WORD, + PUNCTUATION_START, + PUNCTUATION_START_WORD, +) +import probatio + +from homeassistant.auth.permissions.const import CAT_ENTITIES, POLICY_CONTROL +from homeassistant.const import ATTR_ENTITY_ID +from homeassistant.core import HomeAssistant, ServiceCall, SupportsResponse, callback +from homeassistant.exceptions import HomeAssistantError, Unauthorized, UnknownUser +from homeassistant.helpers import config_validation as cv + +from .const import DATA_COMPONENT, DOMAIN, AssistSatelliteEntityFeature +from .entity import AssistSatelliteEntity + + +def has_no_punctuation(value: list[str]) -> list[str]: + """Validate result does not contain punctuation.""" + for sentence in value: + # Exclude {list_references} which may contain punctuation characters. + sentence = _remove_list_references(sentence) + if ( + PUNCTUATION_START.search(sentence) + or PUNCTUATION_END.search(sentence) + or PUNCTUATION_START_WORD.search(sentence) + or PUNCTUATION_END_WORD.search(sentence) + ): + raise probatio.Invalid("sentence should not contain punctuation") + + return value + + +def _remove_list_references(sentence: str) -> str: + """Remove {list_references} from a sentence for linting.""" + return re.sub(r"(? list[str]: + """Validate result can be parsed by hassil.""" + for sentence in value: + try: + parse_sentence(sentence) + except ParseError as err: + raise probatio.Invalid(f"invalid sentence: {err}") from err + return value + + +def has_one_non_empty_item(value: list[str]) -> list[str]: + """Validate result has at least one item.""" + if len(value) < 1: + raise probatio.Invalid("at least one sentence is required") + + for sentence in value: + if not sentence: + raise probatio.Invalid("sentences cannot be empty") + + return value + + +# Validator for media_id fields that accepts both string and media selector format +_media_id_validator = probatio.Any( + cv.string, # Plain string format + probatio.All( + probatio.Schema( + { + probatio.Required("media_content_id"): cv.string, + probatio.Required("media_content_type"): cv.string, + probatio.Remove("metadata"): dict, # Ignore metadata if present + } + ), + # Extract media_content_id from media selector format + lambda x: x["media_content_id"], + ), +) + + +@callback +def async_setup_services(hass: HomeAssistant) -> None: + """Register the Assist satellite services.""" + component = hass.data[DATA_COMPONENT] + + component.async_register_entity_service( + "announce", + probatio.All( + cv.make_entity_service_schema( + { + probatio.Optional("message"): str, + probatio.Optional("media_id"): _media_id_validator, + probatio.Optional("preannounce", default=True): bool, + probatio.Optional("preannounce_media_id"): _media_id_validator, + } + ), + cv.has_at_least_one_key("message", "media_id"), + ), + "async_internal_announce", + [AssistSatelliteEntityFeature.ANNOUNCE], + ) + + component.async_register_entity_service( + "start_conversation", + probatio.All( + cv.make_entity_service_schema( + { + probatio.Optional("start_message"): str, + probatio.Optional("start_media_id"): _media_id_validator, + probatio.Optional("preannounce", default=True): bool, + probatio.Optional("preannounce_media_id"): _media_id_validator, + probatio.Optional("extra_system_prompt"): str, + } + ), + cv.has_at_least_one_key("start_message", "start_media_id"), + ), + "async_internal_start_conversation", + [AssistSatelliteEntityFeature.START_CONVERSATION], + ) + + async def handle_ask_question(call: ServiceCall) -> dict[str, Any]: + """Handle a Show View service call.""" + satellite_entity_id: str = call.data[ATTR_ENTITY_ID] + if call.context.user_id: + user = await hass.auth.async_get_user(call.context.user_id) + if user is None: + raise UnknownUser( + context=call.context, + permission=POLICY_CONTROL, + user_id=call.context.user_id, + ) + if not user.permissions.check_entity(satellite_entity_id, POLICY_CONTROL): + raise Unauthorized( + context=call.context, + permission=POLICY_CONTROL, + user_id=call.context.user_id, + perm_category=CAT_ENTITIES, + ) + + satellite_entity: AssistSatelliteEntity | None = component.get_entity( + satellite_entity_id + ) + if satellite_entity is None: + raise HomeAssistantError( + f"Invalid Assist satellite entity id: {satellite_entity_id}" + ) + + satellite_entity.async_set_context(call.context) + + ask_question_args = { + "question": call.data.get("question"), + "question_media_id": call.data.get("question_media_id"), + "preannounce": call.data.get("preannounce", True), + "answers": call.data.get("answers"), + } + + if preannounce_media_id := call.data.get("preannounce_media_id"): + ask_question_args["preannounce_media_id"] = preannounce_media_id + + answer = await satellite_entity.async_internal_ask_question(**ask_question_args) + + if answer is None: + raise HomeAssistantError("No answer from satellite") + + return asdict(answer) + + hass.services.async_register( + domain=DOMAIN, + service="ask_question", + service_func=handle_ask_question, + schema=probatio.All( + { + probatio.Required(ATTR_ENTITY_ID): cv.entity_domain(DOMAIN), + probatio.Optional("question"): str, + probatio.Optional("question_media_id"): _media_id_validator, + probatio.Optional("preannounce", default=True): bool, + probatio.Optional("preannounce_media_id"): _media_id_validator, + probatio.Optional("answers"): [ + { + probatio.Required("id"): str, + probatio.Required("sentences"): probatio.All( + cv.ensure_list, + [cv.string], + has_one_non_empty_item, + has_no_punctuation, + is_valid_sentence, + ), + } + ], + }, + cv.has_at_least_one_key("question", "question_media_id"), + ), + supports_response=SupportsResponse.ONLY, + )