/
/
1"""
2Home Assistant Plugin for Music Assistant.
3
4The plugin is the core of all communication to/from Home Assistant and
5responsible for maintaining the WebSocket API connection to HA.
6Also, the Music Assistant integration within HA will relay its own api
7communication over the HA api for more flexibility as well as security.
8"""
9
10from __future__ import annotations
11
12import asyncio
13import logging
14import os
15from functools import partial
16from itertools import batched
17from sys import intern
18from types import MappingProxyType
19from typing import TYPE_CHECKING, Any, NamedTuple, TypedDict, cast
20
21from hass_client import HomeAssistantClient
22from hass_client.exceptions import BaseHassClientError
23from hass_client.utils import get_websocket_url
24from music_assistant_models.auth import Scope
25from music_assistant_models.config_entries import ConfigEntry
26from music_assistant_models.enums import (
27 ConfigEntryType,
28 ContentType,
29 EventType,
30 MediaType,
31 ProviderFeature,
32 StreamType,
33)
34from music_assistant_models.errors import (
35 MusicAssistantError,
36 SetupFailedError,
37 UnsupportedFeaturedException,
38)
39from music_assistant_models.media_items.audio_format import AudioFormat
40from music_assistant_models.player_control import PlayerControl
41from music_assistant_models.streamdetails import StreamDetails
42
43from music_assistant.constants import VERBOSE_LOG_LEVEL
44from music_assistant.controllers.cache import use_cache
45from music_assistant.helpers.datetime import iso_from_utc_timestamp
46from music_assistant.helpers.json import SerializableType
47from music_assistant.helpers.util import lock, try_parse_int
48from music_assistant.models.plugin import AIEngine, PluginProvider, TTSEngine
49
50from .constants import (
51 CONF_MUTE_CONTROLS,
52 CONF_POWER_CONTROLS,
53 CONF_VOLUME_CONTROLS,
54 OFF_STATES,
55 MediaPlayerEntityFeature,
56 parse_supported_features,
57)
58from .control_entities import (
59 SEARCH_CONTROL_ENTITIES_LIMIT,
60 ControlEntitySearch,
61 HassControlEntitySearchResult,
62)
63from .helpers import ControlCapabilities, get_control_name, is_entity_id
64
65if TYPE_CHECKING:
66 from collections.abc import Callable, Collection, Mapping
67
68 from aiohttp import ClientResponse, ClientSession
69 from hass_client.models import (
70 Area,
71 CompressedState,
72 Context,
73 Device,
74 Entity,
75 EntityStateEvent,
76 Event,
77 State,
78 )
79 from music_assistant_models.config_entries import ProviderConfig
80 from music_assistant_models.player import PlayerMedia
81 from music_assistant_models.provider import ProviderManifest
82
83 from music_assistant.mass import MusicAssistant
84 from music_assistant.models import ProviderInstanceType
85
86DOMAIN = "hass"
87CONF_URL = "url"
88CONF_AUTH_TOKEN = "token"
89CONF_VERIFY_SSL = "verify_ssl"
90FEATURE_DISCOVERY_TIMEOUT = 30
91STATE_FETCH_TIMEOUT = 30
92STATE_FETCH_BATCH_SIZE = 500
93# window to collect entity registry updates in, so an integration registering a
94# batch of entities results in a single rebuild of the engine lists
95ENGINE_REFRESH_DEBOUNCE = 2
96# window in which repeated device lookups reuse one listing, so a burst of players
97# connecting does not fetch the (unfilterable) device registry once per player
98DEVICE_REGISTRY_CACHE_TTL = 60
99# areas are renamed even less often than devices, and only ever supply a label
100AREA_REGISTRY_CACHE_TTL = 60
101
102SEARCH_CONTROL_ENTITIES_COMMAND = f"{DOMAIN}/search_control_entities"
103
104# Home Assistant entity domains that back the TTS and AI Task features.
105FEATURE_DOMAINS = ("tts", "ai_task")
106FEATURE_DOMAIN_PREFIXES = tuple(f"{domain}." for domain in FEATURE_DOMAINS)
107
108# Entity registry fields a change to which can alter the mirrored registry. Beyond the
109# mirrored fields themselves, disabled_by decides whether an entity is listed at all, and
110# config_entry_id joins them because Home Assistant can clear disabled_by while reporting
111# only the move to the other config entry.
112REGISTRY_FIELDS_AFFECTING_MIRROR = frozenset(
113 {"entity_id", "platform", "device_id", "area_id", "disabled_by", "config_entry_id"}
114)
115
116
117class DeviceMediaPlayerInfo(TypedDict):
118 """Home Assistant correlation info for a device that is natively connected elsewhere."""
119
120 # user-facing device name in HA (name_by_user or name)
121 name: str | None
122 # first enabled media_player entity of the device that supports announcements
123 announce_entity_id: str | None
124
125
126class HassRegistryEntity(NamedTuple):
127 """
128 Home Assistant entity registry entry, limited to the fields Music Assistant uses.
129
130 The entity ID is not a field: entries are always keyed by it.
131 """
132
133 platform: str
134 device_id: str | None
135 # the area the entity is assigned to directly, overriding the one of its device
136 area_id: str | None
137
138
139async def setup(
140 mass: MusicAssistant, manifest: ProviderManifest, config: ProviderConfig
141) -> ProviderInstanceType:
142 """Initialize provider(instance) with given configuration."""
143 return HomeAssistantProvider(mass, manifest, config, set())
144
145
146def _control_config_entries() -> tuple[ConfigEntry, ...]:
147 """Return the config entries holding the entities selected as player controls."""
148 return tuple(
149 ConfigEntry(
150 key=conf_key,
151 type=ConfigEntryType.STRING,
152 multi_value=True,
153 required=True,
154 default_value=[],
155 category="player_controls",
156 )
157 for conf_key in (CONF_POWER_CONTROLS, CONF_VOLUME_CONTROLS, CONF_MUTE_CONTROLS)
158 )
159
160
161class HomeAssistantProvider(PluginProvider):
162 """Home Assistant Plugin for Music Assistant."""
163
164 hass: HomeAssistantClient
165 _listen_task: asyncio.Task[None] | None = None
166 _player_controls: dict[str, PlayerControl] | None = None
167 _unsubscribe_controls: Callable[[], None] | None = None
168 _unsubscribe_entity_registry: Callable[[], None] | None = None
169 _engine_refresh_task: asyncio.Task[None] | None = None
170 _ai_engines: list[AIEngine]
171 _tts_engines: list[TTSEngine]
172 _startup_complete: bool = False
173 _entity_registry: Mapping[str, HassRegistryEntity] | None = None
174 _entity_registry_generation: int = 0
175 _entity_registry_lock: asyncio.Lock
176 _wanted_controls: dict[str, ControlCapabilities] | None = None
177 _control_reconcile_lock: asyncio.Lock
178 _control_entity_search: ControlEntitySearch
179 _unregister_search_command: Callable[[], None] | None = None
180
181 @property
182 def url(self) -> str | None:
183 """Return the configured Home Assistant URL, or None if not configured."""
184 url = self.get_setup_value(CONF_URL)
185 if isinstance(url, str) and url:
186 return url
187 return None
188
189 @property
190 def entity_registry_generation(self) -> int:
191 """Return a counter that changes whenever the mirrored entity registry is dropped."""
192 return self._entity_registry_generation
193
194 async def get_config_entries(self) -> tuple[ConfigEntry, ...]:
195 """
196 Return the (options) config entries for the Home Assistant provider.
197
198 The connection URL and authentication token are collected by the setup flow (see
199 setup_flow.py) unless running as a Home Assistant add-on, where they are fixed; only
200 the player-control and feature options are configurable here.
201 """
202 base_entries: tuple[ConfigEntry, ...]
203 if self.mass.running_as_hass_addon:
204 # on supervisor, we use the internal url
205 # token set to None for auto retrieval
206 base_entries = (
207 ConfigEntry(
208 key=CONF_URL,
209 type=ConfigEntryType.STRING,
210 label=CONF_URL,
211 required=True,
212 default_value="http://supervisor/core/api",
213 value="http://supervisor/core/api",
214 hidden=True,
215 ),
216 ConfigEntry(
217 key=CONF_AUTH_TOKEN,
218 type=ConfigEntryType.STRING,
219 label=CONF_AUTH_TOKEN,
220 required=False,
221 default_value=None,
222 value=None,
223 hidden=True,
224 ),
225 ConfigEntry(
226 key=CONF_VERIFY_SSL,
227 type=ConfigEntryType.BOOLEAN,
228 label=CONF_VERIFY_SSL,
229 required=False,
230 default_value=False,
231 hidden=True,
232 ),
233 )
234 else:
235 # url/token/verify_ssl are collected by the setup flow instead (see setup_flow.py)
236 base_entries = ()
237
238 return (*base_entries, *_control_config_entries())
239
240 async def handle_async_init(self) -> None:
241 """Handle async initialization of the plugin."""
242 if self._listen_task and not self._listen_task.done():
243 msg = "Home Assistant listener is already running"
244 raise SetupFailedError(msg)
245 self._startup_complete = False
246 self._player_controls = {}
247 self._wanted_controls = None
248 self._control_reconcile_lock = asyncio.Lock()
249 self._ai_engines = []
250 self._tts_engines = []
251 url = get_websocket_url(cast("str", self.get_setup_value(CONF_URL)))
252 token = self.get_setup_value(CONF_AUTH_TOKEN)
253 logging.getLogger("hass_client").setLevel(self.logger.level + 10)
254 ssl = bool(self.get_setup_value(CONF_VERIFY_SSL, True))
255 http_session = self.mass.http_session if ssl else self.mass.http_session_no_ssl
256 self.hass = HomeAssistantClient(url, token, http_session)
257 self._entity_registry = None
258 self._entity_registry_lock = asyncio.Lock()
259 self._control_entity_search = ControlEntitySearch(self)
260 # registering here rather than in loaded_in_mass pairs the command with the teardown
261 # in _disconnect_hass, so a reload can never leave it registered twice
262 self._unregister_search_command = self.mass.register_api_command(
263 SEARCH_CONTROL_ENTITIES_COMMAND,
264 self.search_control_entities,
265 required_scope=Scope.CONFIG_PROVIDERS_READ,
266 )
267 try:
268 await self.hass.connect()
269 except BaseHassClientError as err:
270 await self._cleanup_failed_init()
271 err_msg = str(err) or err.__class__.__name__
272 raise SetupFailedError(err_msg) from err
273 self._listen_task = self.mass.create_task(self._hass_listener())
274 try:
275 # the registry subscription must be live before the first registry read, so no
276 # registry change can slip through unnoticed; _disconnect_hass tears the
277 # subscription down again on the failure paths below
278 await self._subscribe_entity_registry()
279 await self._resolve_startup_features()
280 except asyncio.CancelledError:
281 await self._cleanup_failed_init()
282 raise
283 except BaseHassClientError as err:
284 await self._cleanup_failed_init()
285 err_msg = str(err) or err.__class__.__name__
286 raise SetupFailedError(err_msg) from err
287 except Exception:
288 await self._cleanup_failed_init()
289 raise
290
291 async def loaded_in_mass(self) -> None:
292 """Call after the provider has been loaded."""
293 await self._register_player_controls()
294
295 async def unload(self, is_removed: bool = False) -> None:
296 """
297 Handle unload/close of the provider.
298
299 Called when provider is deregistered (e.g. MA exiting or config reloading).
300 """
301 # unregister all player controls
302 if self._player_controls:
303 for entity_id in self._player_controls:
304 self.mass.players.remove_player_control(entity_id)
305 self._startup_complete = False
306 await self._disconnect_hass()
307
308 async def update_config(self, config: ProviderConfig, changed_keys: set[str]) -> None:
309 """
310 Handle logic when the config is updated.
311
312 A change limited to the player control selection is applied in place, so adding or
313 removing a control does not drop and re-establish the Home Assistant connection.
314 Any other change reloads the provider as usual.
315
316 Raises when the in place update fails (because Home Assistant is unreachable, for
317 example); the controls are then left as they were and a later update retries.
318
319 :param config: The updated provider config.
320 :param changed_keys: The keys that changed in the given config.
321 """
322 control_keys = {
323 f"values/{conf_key}"
324 for conf_key in (CONF_POWER_CONTROLS, CONF_VOLUME_CONTROLS, CONF_MUTE_CONTROLS)
325 }
326 if not changed_keys or not changed_keys <= control_keys:
327 await super().update_config(config, changed_keys)
328 return
329 # store the new config before reconciling: the control lists are read back from it
330 self.config = config
331 await self._register_player_controls()
332
333 async def get_diagnostics(self) -> dict[str, SerializableType]:
334 """Return diagnostics info for this provider to include in diagnostics reports."""
335 return {
336 "connected": self.hass.connected,
337 "ha_version": self.hass.version,
338 "listener_active": self._listen_task is not None and not self._listen_task.done(),
339 "player_controls": len(self._player_controls) if self._player_controls else 0,
340 }
341
342 async def get_entity_registry(self) -> Mapping[str, HassRegistryEntity]:
343 """
344 Return the Home Assistant entity registry, keyed by entity ID.
345
346 Entities that are disabled in Home Assistant are absent from the result, and so are
347 entities without a unique ID: those are not part of Home Assistant's registry at all.
348
349 The result is shared between all callers and is read-only: both the mapping and
350 its entries reject writes.
351 """
352 if (registry := self._entity_registry) is not None:
353 return registry
354 async with self._entity_registry_lock:
355 if (registry := self._entity_registry) is None:
356 generation = self._entity_registry_generation
357 registry = await self._fetch_entity_registry()
358 # a registry change while the fetch was in flight leaves the listing stale
359 # on arrival, so serve it to this caller but keep it out of the cache
360 if generation == self._entity_registry_generation:
361 self._entity_registry = registry
362 return registry
363
364 async def get_entity_registry_entries(self, entity_ids: Collection[str]) -> dict[str, Entity]:
365 """
366 Return the full Home Assistant entity registry entries of the given entities.
367
368 :param entity_ids: The entity IDs to look up.
369 :return: The registry entries keyed by entity ID; entities unknown to
370 Home Assistant are absent from the result.
371 """
372 if not entity_ids:
373 return {}
374 result = cast(
375 "dict[str, Entity | None]",
376 await self.hass.send_command(
377 "config/entity_registry/get_entries", entity_ids=list(entity_ids)
378 ),
379 )
380 return {entity_id: entry for entity_id, entry in result.items() if entry is not None}
381
382 async def get_device_registry(self) -> dict[str, Device]:
383 """
384 Return the Home Assistant device registry, keyed by device ID.
385
386 Home Assistant offers no abbreviated variant of the device registry listing, so the
387 entries carry all of their fields. The listing is reused for a short while, so a
388 device change may take up to DEVICE_REGISTRY_CACHE_TTL seconds to be reflected.
389 """
390 return await self._fetch_device_registry()
391
392 async def get_area_registry(self) -> dict[str, Area]:
393 """
394 Return the Home Assistant area registry, keyed by area ID.
395
396 The listing is reused for a short while, so an area change may take up to
397 AREA_REGISTRY_CACHE_TTL seconds to be reflected.
398 """
399 return await self._fetch_area_registry()
400
401 async def search_control_entities(
402 self,
403 search: str | None = None,
404 control_type: str | None = None,
405 limit: int = SEARCH_CONTROL_ENTITIES_LIMIT,
406 ) -> HassControlEntitySearchResult:
407 """
408 Search the Home Assistant entities that can be used as a player control.
409
410 Music Assistant's own players are never part of the result. Consecutive searches are
411 served from a short lived cache that an entity registry change drops right away, so a
412 newly added or removed entity shows up immediately, while a device or area rename can
413 lag by up to a minute.
414
415 :param search: Text to match, case insensitively, against the entity ID, the entity
416 name, its device name and its area name. Every whitespace separated word must
417 match one of those fields, though not necessarily the same one. All eligible
418 entities match when omitted.
419 :param control_type: Restrict the result to entities that can serve this control role,
420 given as one of the provider's control config keys (``power_controls``,
421 ``volume_controls`` or ``mute_controls``). All roles are returned when omitted.
422 :param limit: Maximum number of entities (not groups) to return, itself capped at
423 ``SEARCH_CONTROL_ENTITIES_MAX_LIMIT``.
424 :return: The matching entities grouped by the device and area they belong to, ordered
425 by area, device and entity name, plus a flag telling whether matches were left out
426 to honor the limit.
427 """
428 return await self._control_entity_search.search(search, control_type, limit)
429
430 async def get_media_player_device_infos(
431 self,
432 mac_addresses: Collection[str],
433 platform: str,
434 ) -> dict[str, DeviceMediaPlayerInfo]:
435 """
436 Correlate devices (by MAC address) to their HA name and media_player entity.
437
438 Used for devices that are natively connected to Music Assistant but also
439 present in Home Assistant, to pick up their HA device name and their
440 (announcement-capable) media_player entity.
441
442 :param mac_addresses: Device MAC addresses to look up (case-insensitive).
443 :param platform: The HA integration domain the media_player entities must belong to.
444 :return: Correlation info keyed by lowercased MAC address; devices unknown
445 to Home Assistant are absent from the result.
446 """
447 wanted_macs = {mac.lower() for mac in mac_addresses}
448 if not wanted_macs:
449 return {}
450 device_registry = await self.get_device_registry()
451 device_by_mac: dict[str, Device] = {
452 connection[1].lower(): device
453 for device in device_registry.values()
454 for connection in device.get("connections", [])
455 if len(connection) == 2
456 and connection[0] == "mac"
457 and connection[1].lower() in wanted_macs
458 }
459 if not device_by_mac:
460 return {}
461 media_players_by_device: dict[str, list[str]] = {}
462 for entity_id, entry in (await self.get_entity_registry()).items():
463 if (
464 entry.platform == platform
465 and entity_id.startswith("media_player.")
466 and (device_id := entry.device_id)
467 ):
468 media_players_by_device.setdefault(device_id, []).append(entity_id)
469 candidates_by_mac = {
470 mac: media_players_by_device.get(device["id"], [])
471 for mac, device in device_by_mac.items()
472 }
473 states = {
474 state["entity_id"]: state
475 for state in await self.get_states(
476 entity_ids=[
477 entity_id
478 for entity_ids in candidates_by_mac.values()
479 for entity_id in entity_ids
480 ]
481 )
482 }
483
484 def _supports_announce(entity_id: str) -> bool:
485 if (state := states.get(entity_id)) is None:
486 return False
487 supported_features = parse_supported_features(
488 state["attributes"].get("supported_features"), entity_id, self.logger
489 )
490 return MediaPlayerEntityFeature.MEDIA_ANNOUNCE in supported_features
491
492 return {
493 mac: DeviceMediaPlayerInfo(
494 name=device["name_by_user"] or device["name"],
495 announce_entity_id=next(
496 (
497 entity_id
498 for entity_id in candidates_by_mac[mac]
499 if _supports_announce(entity_id)
500 ),
501 None,
502 ),
503 )
504 for mac, device in device_by_mac.items()
505 }
506
507 async def get_user_details(self, ha_user_id: str) -> tuple[str | None, str | None, str | None]:
508 """
509 Get user username, display name and avatar URL from Home Assistant.
510
511 Looks up the user in config/auth/list for username, and the person entity
512 for display name and picture URL.
513
514 :param ha_user_id: Home Assistant user ID.
515 :return: Tuple of (username, display_name, avatar_url) or all None if not found.
516 """
517 try:
518 username: str | None = None
519 display_name: str | None = None
520 avatar_url: str | None = None
521
522 # Get username from config/auth/list (admin endpoint, we have admin access)
523 try:
524 users = await self.hass.send_command("config/auth/list")
525 for user in users or []:
526 if user.get("id") == ha_user_id:
527 username = user.get("username")
528 # Also get name as fallback display name
529 if not display_name:
530 display_name = user.get("name")
531 break
532 except Exception as err:
533 self.logger.log(VERBOSE_LOG_LEVEL, "Failed to get HA user list: %s", err)
534
535 # Get external URL for building avatar URL
536 ha_url: str | None = None
537 try:
538 network_urls = await self.hass.send_command("network/url")
539 if network_urls:
540 ha_url = network_urls.get("external") or network_urls.get("internal")
541 except Exception as err:
542 self.logger.log(VERBOSE_LOG_LEVEL, "Failed to get HA network URLs: %s", err)
543
544 # Find person linked to this HA user ID for display name and avatar
545 try:
546 persons = await self.hass.send_command("person/list")
547 # person/list returns {storage: [...], config: [...]}
548 all_persons = (persons.get("storage") or []) + (persons.get("config") or [])
549 for person in all_persons:
550 if person.get("user_id") == ha_user_id:
551 # Person name takes priority for display name
552 if person_name := person.get("name"):
553 display_name = person_name
554 if (person_picture := person.get("picture")) and ha_url:
555 avatar_url = f"{ha_url.rstrip('/')}{person_picture}"
556 break
557 except Exception as err:
558 self.logger.log(VERBOSE_LOG_LEVEL, "Failed to get HA person details: %s", err)
559
560 self.logger.log(
561 VERBOSE_LOG_LEVEL,
562 "get_user_details for %s: username=%s, display_name=%s, avatar_url=%s",
563 ha_user_id,
564 username,
565 display_name,
566 avatar_url,
567 )
568 return username, display_name, avatar_url
569 except Exception as err:
570 self.logger.warning("Failed to get HA user details: %s", err)
571 return None, None, None
572
573 async def get_states(
574 self,
575 *,
576 entity_ids: list[str] | None = None,
577 domains: Collection[str] | None = None,
578 ) -> list[State]:
579 """
580 Return the current Home Assistant state for the requested entities.
581
582 Provide explicit entity IDs and/or a set of domains; only those entities
583 are fetched.
584
585 :param entity_ids: Explicit entity IDs to fetch the current state for.
586 :param domains: Entity domains whose entities should be fetched.
587 """
588 ids: set[str] = set(entity_ids or ())
589 if domains:
590 # resolve domains to entity_ids via the registry, which is far smaller
591 # than a full state dump (it carries no attributes)
592 registry = await self.get_entity_registry()
593 ids.update(entity_id for entity_id in registry if entity_id.split(".", 1)[0] in domains)
594 if not ids:
595 return []
596 states: list[State] = []
597 async with asyncio.timeout(STATE_FETCH_TIMEOUT):
598 # exceeding hass_client's 16MB websocket message limit drops the entire
599 # connection, so bound the state dump by construction and fetch in batches
600 for batch in batched(sorted(ids), STATE_FETCH_BATCH_SIZE, strict=False):
601 states.extend(await self._fetch_states(list(batch)))
602 return states
603
604 async def resolve_image(self, path: str) -> bytes:
605 """Resolve an image from an image path."""
606 ha_url, headers, http_session = self._get_ha_http()
607 async with http_session.get(f"{ha_url}{path}", headers=headers) as response:
608 response.raise_for_status()
609 return await response.read()
610
611 async def get_ai_engines(self) -> list[AIEngine]:
612 """Return the Home Assistant AI Task entities as AI engines."""
613 return self._ai_engines
614
615 async def get_tts_engines(self) -> list[TTSEngine]:
616 """Return the Home Assistant TTS entities as TTS engines."""
617 return self._tts_engines
618
619 async def ai_query(self, query: str, engine_id: str | None = None) -> str:
620 """Handle an AI query via Home Assistant's ai_task service."""
621 entity_id = engine_id or next((engine.id for engine in self._ai_engines), None)
622 if entity_id is None:
623 raise UnsupportedFeaturedException("AI Task entity is not available")
624 result = await self.hass.send_command(
625 "call_service",
626 domain="ai_task",
627 service="generate_data",
628 service_data={
629 "task_name": "music_assistant",
630 "instructions": query,
631 "entity_id": entity_id,
632 },
633 return_response=True,
634 )
635 response = result.get("response", {}) if isinstance(result, dict) else {}
636 data = response.get("data") if isinstance(response, dict) else None
637 if not data:
638 msg = f"AI Task returned no data in response: {result}"
639 raise MusicAssistantError(msg)
640 return str(data)
641
642 async def play_announcement_on_entity(self, entity_id: str, announcement: PlayerMedia) -> None:
643 """
644 Play an announcement on a Home Assistant media_player entity.
645
646 Uses Home Assistant's announce feature, so the entity's integration ducks
647 or pauses any running playback and resumes it afterwards. Returns once the
648 announcement has finished playing (approximated by its duration).
649
650 :param entity_id: The media_player entity to play the announcement on.
651 :param announcement: The announcement to play.
652 """
653 await self.hass.call_service(
654 domain="media_player",
655 service="play_media",
656 service_data={
657 "media_content_id": announcement.uri,
658 "media_content_type": "music",
659 "announce": True,
660 },
661 target={"entity_id": entity_id},
662 )
663 # Wait until the announcement is finished playing so callers can play
664 # announcements in a sequence; HA gives no completion signal for announcements.
665 duration = await self.mass.streams.get_announcement_duration(announcement)
666 await asyncio.sleep(duration or 5)
667
668 async def get_tts_message(
669 self,
670 message: str,
671 language: str | None = None,
672 engine_id: str | None = None,
673 options: dict[str, Any] | None = None,
674 ) -> StreamDetails:
675 """Handle text-to-speech via Home Assistant's REST API."""
676 entity_id = engine_id or next((engine.id for engine in self._tts_engines), None)
677 if entity_id is None:
678 raise UnsupportedFeaturedException("TTS entity is not available")
679 ha_url, headers, http_session = self._get_ha_http()
680 # the tts_get_url payload field is called engine_id but takes a tts entity_id
681 payload: dict[str, Any] = {"engine_id": entity_id, "message": message}
682 if language:
683 payload["language"] = language
684 if options:
685 payload["options"] = options
686 async with http_session.post(
687 f"{ha_url}/api/tts_get_url", headers=headers, json=payload
688 ) as response:
689 await self._raise_for_tts_error(response)
690 data = await response.json()
691 url = str(data["url"])
692 return StreamDetails(
693 provider=self.instance_id,
694 item_id=url,
695 audio_format=AudioFormat(content_type=ContentType.MP3),
696 media_type=MediaType.SOUND_EFFECT,
697 stream_type=StreamType.HTTP,
698 path=url,
699 )
700
701 async def _hass_listener(self) -> None:
702 """Start listening on the HA websockets."""
703 try:
704 # start listening will block until the connection is lost/closed
705 await self.hass.start_listening()
706 except BaseHassClientError as err:
707 self.logger.warning("Connection to HA lost due to error: %s", err)
708 if not self._startup_complete:
709 return
710 self.logger.info("Connection to HA lost. Connection will be automatically retried later.")
711 # schedule a reload of the provider, armed under the load path's task id so any
712 # (re)load starting before it fires cancels it
713 self.available = False
714 self.mass.call_later(
715 5,
716 self.mass.load_provider,
717 self.instance_id,
718 allow_retry=True,
719 task_id=f"load_provider_{self.instance_id}",
720 )
721
722 def _on_entity_state_update(self, event: EntityStateEvent) -> None:
723 """Handle Entity State event."""
724 if entity_additions := event.get("a"):
725 for entity_id, state in entity_additions.items():
726 self._update_control_from_state_msg(entity_id, state)
727 if entity_changes := event.get("c"):
728 for entity_id, state_diff in entity_changes.items():
729 if "+" not in state_diff:
730 continue
731 self._update_control_from_state_msg(entity_id, state_diff["+"])
732
733 async def _register_player_controls(self) -> None:
734 """Bring the registered player controls in line with the current configuration."""
735 assert self._player_controls is not None # for type checking
736 # the wanted selection is determined inside the lock, so a reconcile that had to
737 # wait for another one cannot apply a selection that was already superseded
738 async with self._control_reconcile_lock:
739 power_controls = self._selected_control_entities(CONF_POWER_CONTROLS)
740 mute_controls = self._selected_control_entities(CONF_MUTE_CONTROLS)
741 volume_controls = self._selected_control_entities(CONF_VOLUME_CONTROLS)
742 wanted_controls: dict[str, ControlCapabilities] = {
743 entity_id: ControlCapabilities(
744 power=entity_id in power_controls,
745 volume=entity_id in volume_controls,
746 mute=entity_id in mute_controls,
747 )
748 for entity_id in (*power_controls, *mute_controls, *volume_controls)
749 }
750 if wanted_controls == self._wanted_controls:
751 # the selection is unchanged, so there is no need to consult Home Assistant
752 return
753 hass_states = {
754 state["entity_id"]: state
755 for state in await self.get_states(entity_ids=list(wanted_controls))
756 }
757 for entity_id in set(self._player_controls) - set(wanted_controls):
758 del self._player_controls[entity_id]
759 self.mass.players.remove_player_control(entity_id)
760 for entity_id, capabilities in wanted_controls.items():
761 control = self._create_player_control(
762 entity_id, hass_states.get(entity_id), capabilities
763 )
764 self._player_controls[entity_id] = control
765 await self.mass.players.register_or_update_player_control(control)
766 await self._subscribe_control_states()
767 self._wanted_controls = wanted_controls
768
769 def _selected_control_entities(self, conf_key: str) -> list[str]:
770 """
771 Return the entity IDs selected in the given player control setting.
772
773 :param conf_key: The control config key to read the selection from.
774 """
775 entity_ids: list[str] = []
776 for value in cast("list[str]", self.config.get_value(conf_key)):
777 if is_entity_id(value):
778 entity_ids.append(value)
779 continue
780 # Home Assistant rejects an entire state fetch or subscription over a single
781 # value that is not an entity ID, so a leftover selection would otherwise
782 # take down every control of this provider
783 self.logger.warning(
784 "Ignoring %r in the %s setting: it is not a Home Assistant entity ID",
785 value,
786 conf_key,
787 )
788 return entity_ids
789
790 def _create_player_control(
791 self,
792 entity_id: str,
793 hass_state: State | None,
794 capabilities: ControlCapabilities,
795 ) -> PlayerControl:
796 """
797 Return a ready to use PlayerControl for a Home Assistant entity.
798
799 :param entity_id: The entity to base the control on.
800 :param hass_state: The entity's current state, if known.
801 :param capabilities: The control roles the entity should serve.
802 """
803 entity_platform = entity_id.split(".", maxsplit=1)[0]
804 control = PlayerControl(
805 id=entity_id,
806 provider=self.instance_id,
807 name=get_control_name(entity_id, hass_state),
808 )
809 if capabilities.power:
810 control.supports_power = True
811 control.power_state = hass_state["state"] not in OFF_STATES if hass_state else False
812 control.power_on = partial(self._handle_player_control_power_on, entity_id)
813 control.power_off = partial(self._handle_player_control_power_off, entity_id)
814 if capabilities.volume:
815 control.supports_volume = True
816 if not hass_state:
817 control.volume_level = 0
818 elif entity_platform == "media_player":
819 control.volume_level = int(hass_state["attributes"].get("volume_level", 0) * 100)
820 else:
821 control.volume_level = try_parse_int(hass_state["state"]) or 0
822 control.volume_set = partial(self._handle_player_control_volume_set, entity_id)
823 if capabilities.mute:
824 control.supports_mute = True
825 if not hass_state:
826 control.volume_muted = False
827 elif entity_platform == "media_player":
828 control.volume_muted = bool(hass_state["attributes"].get("is_volume_muted"))
829 else:
830 control.volume_muted = hass_state["state"] not in OFF_STATES
831 control.mute_set = partial(self._handle_player_control_mute_set, entity_id)
832 return control
833
834 async def _subscribe_control_states(self) -> None:
835 """Subscribe to the Home Assistant state of all currently tracked controls."""
836 assert self._player_controls is not None # for type checking
837 # the earlier subscription is only released once the new one is live, so a failure
838 # to subscribe leaves the controls watched by the subscription they already had
839 previous_unsubscribe = self._unsubscribe_controls
840 self._unsubscribe_controls = await self.hass.subscribe_entities(
841 self._on_entity_state_update, list(self._player_controls)
842 )
843 if previous_unsubscribe:
844 previous_unsubscribe()
845
846 async def _handle_player_control_power_on(self, entity_id: str) -> None:
847 """Handle powering on the playercontrol."""
848 await self.hass.call_service(
849 domain="homeassistant",
850 service="turn_on",
851 target={"entity_id": entity_id},
852 )
853
854 async def _handle_player_control_power_off(self, entity_id: str) -> None:
855 """Handle powering off the playercontrol."""
856 await self.hass.call_service(
857 domain="homeassistant",
858 service="turn_off",
859 target={"entity_id": entity_id},
860 )
861
862 async def _handle_player_control_mute_set(self, entity_id: str, muted: bool) -> None:
863 """Handle muting the playercontrol."""
864 if entity_id.startswith("media_player."):
865 await self.hass.call_service(
866 domain="media_player",
867 service="volume_mute",
868 service_data={"is_volume_muted": muted},
869 target={"entity_id": entity_id},
870 )
871 else:
872 await self.hass.call_service(
873 domain="homeassistant",
874 service="turn_off" if muted else "turn_on",
875 target={"entity_id": entity_id},
876 )
877
878 async def _handle_player_control_volume_set(self, entity_id: str, volume_level: int) -> None:
879 """Handle setting volume on the playercontrol."""
880 domain = entity_id.split(".", 1)[0]
881
882 if domain == "media_player":
883 await self.hass.call_service(
884 domain=domain,
885 service="volume_set",
886 service_data={"volume_level": volume_level / 100},
887 target={"entity_id": entity_id},
888 )
889 return
890
891 # At this point, `set_value` will work for both `number` or `input_number`
892 await self.hass.call_service(
893 domain=domain,
894 service="set_value",
895 target={"entity_id": entity_id},
896 service_data={"value": volume_level},
897 )
898
899 def _update_control_from_state_msg(self, entity_id: str, state: CompressedState) -> None:
900 """Update PlayerControl from state(update) message."""
901 if self._player_controls is None:
902 return
903 if not (player_control := self._player_controls.get(entity_id)):
904 return
905 entity_platform = entity_id.split(".", maxsplit=1)[0]
906 if "s" in state:
907 # state changed
908 if player_control.supports_power:
909 player_control.power_state = state["s"] not in OFF_STATES
910 if player_control.supports_mute and entity_platform != "media_player":
911 player_control.volume_muted = state["s"] not in OFF_STATES
912 if player_control.supports_volume and entity_platform != "media_player":
913 player_control.volume_level = try_parse_int(state["s"]) or 0
914 if "a" in state and (attributes := state["a"]):
915 if player_control.supports_volume and "volume_level" in attributes:
916 player_control.volume_level = int(attributes.get("volume_level", 0) * 100)
917 if player_control.supports_mute and "is_volume_muted" in attributes:
918 player_control.volume_muted = bool(attributes.get("is_volume_muted"))
919 self.mass.players.update_player_control(entity_id)
920
921 async def _fetch_states(self, entity_ids: list[str]) -> list[State]:
922 """
923 Return the current Home Assistant state of the given entities.
924
925 :param entity_ids: The entity IDs to fetch the current state for.
926 :return: The states of the requested entities; entities that currently have
927 no state are absent from the result.
928 """
929 initial_states: asyncio.Future[dict[str, CompressedState]]
930 initial_states = asyncio.get_running_loop().create_future()
931
932 def _on_initial_states(event: EntityStateEvent) -> None:
933 # only the first message of a subscription carries the full state under "a";
934 # a state change racing in ahead of it must not resolve the fetch
935 if (added := event.get("a")) is not None and not initial_states.done():
936 initial_states.set_result(added)
937
938 unsubscribe = await self.hass.subscribe_entities(_on_initial_states, entity_ids)
939 try:
940 compressed_states = await initial_states
941 finally:
942 unsubscribe()
943 return [
944 _decompress_state(entity_id, compressed_state)
945 for entity_id, compressed_state in compressed_states.items()
946 ]
947
948 def _get_ha_http(self) -> tuple[str, dict[str, str], ClientSession]:
949 """Return HA base URL (without trailing /api), auth headers, and the HTTP session."""
950 ha_url = cast("str", self.get_setup_value(CONF_URL)).rstrip("/")
951 ha_url = ha_url.removesuffix("/api")
952 token = self.get_setup_value(CONF_AUTH_TOKEN) or os.environ.get("HASSIO_TOKEN")
953 headers = {"Authorization": f"Bearer {token}"} if token else {}
954 ssl = bool(self.get_setup_value(CONF_VERIFY_SSL, True))
955 http_session = self.mass.http_session if ssl else self.mass.http_session_no_ssl
956 return ha_url, headers, http_session
957
958 async def _raise_for_tts_error(self, response: ClientResponse) -> None:
959 """Raise for a failed tts_get_url response without masking a language rejection."""
960 if response.ok:
961 return
962 try:
963 # content_type=None so an error body served as text/plain still parses
964 body = await response.json(content_type=None)
965 except ValueError:
966 body = None
967 error_message = body.get("error") if isinstance(body, dict) else None
968 # HA returns the same 400 for a rejected option and an unsupported language
969 if isinstance(error_message, str) and "Invalid options found" in error_message:
970 raise MusicAssistantError(error_message)
971 response.raise_for_status()
972
973 async def _disconnect_hass(self) -> None:
974 """Stop listening for Home Assistant events and disconnect the client."""
975 if unregister := self._unregister_search_command:
976 self._unregister_search_command = None
977 unregister()
978 self._control_entity_search.close()
979 if unsubscribe := self._unsubscribe_controls:
980 self._unsubscribe_controls = None
981 unsubscribe()
982 if unsubscribe := self._unsubscribe_entity_registry:
983 self._unsubscribe_entity_registry = None
984 unsubscribe()
985 if refresh_task := self._engine_refresh_task:
986 self._engine_refresh_task = None
987 refresh_task.cancel()
988 if listen_task := self._listen_task:
989 self._listen_task = None
990 if not listen_task.done():
991 listen_task.cancel()
992 try:
993 await listen_task
994 except asyncio.CancelledError:
995 pass
996 except Exception as err:
997 self.logger.warning("Home Assistant listener stopped with error: %s", err)
998 await self.hass.disconnect()
999
1000 async def _cleanup_failed_init(self) -> None:
1001 """Clean up the Home Assistant connection after initialization fails."""
1002 try:
1003 await self._disconnect_hass()
1004 except Exception as err:
1005 self.logger.warning("Failed to disconnect from Home Assistant: %s", err)
1006
1007 async def _resolve_startup_features(self) -> None:
1008 """Resolve Home Assistant features while the listener remains active."""
1009 assert self._listen_task is not None
1010 feature_task = asyncio.create_task(self._refresh_engines())
1011 try:
1012 try:
1013 async with asyncio.timeout(FEATURE_DISCOVERY_TIMEOUT):
1014 await asyncio.wait(
1015 {feature_task, self._listen_task},
1016 return_when=asyncio.FIRST_COMPLETED,
1017 )
1018 except TimeoutError as err:
1019 msg = "Timed out while resolving Home Assistant feature entities"
1020 raise SetupFailedError(msg) from err
1021 if not feature_task.done():
1022 msg = "Home Assistant listener stopped during startup"
1023 raise SetupFailedError(msg)
1024 if feature_task.cancelled():
1025 if self._listen_task.done():
1026 msg = "Home Assistant listener stopped during startup"
1027 else:
1028 msg = "Home Assistant feature resolution was cancelled"
1029 raise SetupFailedError(msg)
1030 await feature_task
1031 if self._listen_task.done():
1032 msg = "Home Assistant listener stopped during startup"
1033 raise SetupFailedError(msg)
1034 self._startup_complete = True
1035 finally:
1036 if not feature_task.done():
1037 feature_task.cancel()
1038 await asyncio.gather(feature_task, return_exceptions=True)
1039
1040 async def _refresh_engines(self) -> None:
1041 """Rebuild the TTS/AI engine lists from the Home Assistant feature entities."""
1042 tts_engines: list[TTSEngine] = []
1043 ai_engines: list[AIEngine] = []
1044 for state in await self.get_states(domains=FEATURE_DOMAINS):
1045 entity_id = state["entity_id"]
1046 entity_platform = entity_id.split(".", 1)[0]
1047 if friendly_name := state["attributes"].get("friendly_name"):
1048 name = f"{friendly_name} ({entity_id})"
1049 else:
1050 name = entity_id
1051 if entity_platform == "tts":
1052 tts_engines.append(TTSEngine(id=entity_id, name=name, provider=self))
1053 elif entity_platform == "ai_task":
1054 ai_engines.append(AIEngine(id=entity_id, name=name, provider=self))
1055 tts_engines.sort(key=lambda engine: engine.name)
1056 ai_engines.sort(key=lambda engine: engine.name)
1057 changed = (self._tts_engines, self._ai_engines) != (tts_engines, ai_engines)
1058 self._tts_engines = tts_engines
1059 self._ai_engines = ai_engines
1060 self._supported_features.discard(ProviderFeature.TTS)
1061 self._supported_features.discard(ProviderFeature.AI_QUERY)
1062 if tts_engines:
1063 self._supported_features.add(ProviderFeature.TTS)
1064 if ai_engines:
1065 self._supported_features.add(ProviderFeature.AI_QUERY)
1066 # the entities can come and go without this provider (un)loading, so tell the
1067 # consumers of our engines that their selection may need re-evaluating. They read
1068 # the lists straight from their handler, so this has to stay below the assignments.
1069 # during startup the load itself signals once we are done.
1070 if changed and self._startup_complete:
1071 self.mass.signal_event(EventType.PROVIDERS_UPDATED, data=self.mass.get_providers())
1072
1073 async def _subscribe_entity_registry(self) -> None:
1074 """Watch the Home Assistant entity registry to keep the engine lists up to date."""
1075 # register for entity registry updates, replacing any earlier subscription
1076 if unsubscribe := self._unsubscribe_entity_registry:
1077 self._unsubscribe_entity_registry = None
1078 unsubscribe()
1079 self._unsubscribe_entity_registry = await self.hass.subscribe_events(
1080 self._on_entity_registry_update, "entity_registry_updated"
1081 )
1082
1083 def _on_entity_registry_update(self, event: Event) -> None:
1084 """Handle an entity registry update event."""
1085 data = event["data"]
1086 if _affects_mirrored_registry(data):
1087 self._entity_registry = None
1088 self._entity_registry_generation += 1
1089 elif self.logger.isEnabledFor(VERBOSE_LOG_LEVEL):
1090 # the kept mirror rests on which fields Home Assistant reports, so leave a
1091 # trail that tells a stale entity listing apart from a lookup that never ran
1092 self.logger.log(
1093 VERBOSE_LOG_LEVEL,
1094 "Keeping the mirrored entity registry, %s changed on %s",
1095 ", ".join(data["changes"]),
1096 data.get("entity_id", "?"),
1097 )
1098 entity_id = data.get("entity_id", "")
1099 if not entity_id.startswith(FEATURE_DOMAIN_PREFIXES):
1100 return
1101 self._schedule_engine_refresh()
1102
1103 def _schedule_engine_refresh(self) -> None:
1104 """(Re)schedule the debounced rebuild of the engine lists."""
1105 if refresh_task := self._engine_refresh_task:
1106 self._engine_refresh_task = None
1107 refresh_task.cancel()
1108 self._engine_refresh_task = self.mass.create_task(self._delayed_engine_refresh())
1109
1110 async def _delayed_engine_refresh(self) -> None:
1111 """Rebuild the engine lists once the debounce window has passed."""
1112 await asyncio.sleep(ENGINE_REFRESH_DEBOUNCE)
1113 try:
1114 await self._refresh_engines()
1115 except Exception as err:
1116 self.logger.warning("Failed to refresh Home Assistant engines: %s", err)
1117
1118 # unlike _fetch_device_registry, this listing is mirrored for the lifetime of the
1119 # connection rather than kept behind a TTL: it runs to several megabytes on a large
1120 # setup, and Home Assistant announces every change, so the mirror is both the cheaper
1121 # and the more accurate option
1122 async def _fetch_entity_registry(self) -> Mapping[str, HassRegistryEntity]:
1123 """Fetch the entity registry from Home Assistant, keyed by entity ID."""
1124 # the display variant of the registry listing carries abbreviated keys and only the
1125 # fields the Home Assistant frontend needs, making it several times smaller
1126 result = cast(
1127 "dict[str, Any]",
1128 await self.hass.send_command("config/entity_registry/list_for_display"),
1129 )
1130 # the listing repeats a handful of platform names, one device id per device and one
1131 # area id per area over all of its entities, so hold on to a single string object
1132 # per distinct value
1133 device_ids: dict[str, str] = {}
1134 area_ids: dict[str, str] = {}
1135 registry: dict[str, HassRegistryEntity] = {}
1136 for entry in result["entities"]:
1137 if (device_id := entry.get("di")) is not None:
1138 device_id = device_ids.setdefault(device_id, device_id)
1139 if (area_id := entry.get("ai")) is not None:
1140 area_id = area_ids.setdefault(area_id, area_id)
1141 registry[entry["ei"]] = HassRegistryEntity(
1142 platform=intern(entry["pl"]),
1143 device_id=device_id,
1144 area_id=area_id,
1145 )
1146 return MappingProxyType(registry)
1147
1148 # the lock sits outside the cache to keep a burst of lookups from fetching once per
1149 # caller. use_cache stores in the background, so the callers that reach a still-cold
1150 # cache can overlap: the burst costs a couple of fetches instead of one per player
1151 @lock
1152 @use_cache(expiration=DEVICE_REGISTRY_CACHE_TTL)
1153 async def _fetch_device_registry(self) -> dict[str, Any]:
1154 """Fetch the device registry from Home Assistant, keyed by device ID."""
1155 # use_cache rebuilds the cached value from this return annotation, which rules out
1156 # the Device TypedDict; get_device_registry restores the type for callers
1157 return {device["id"]: device for device in await self.hass.get_device_registry()}
1158
1159 @lock
1160 @use_cache(expiration=AREA_REGISTRY_CACHE_TTL)
1161 async def _fetch_area_registry(self) -> dict[str, Any]:
1162 """Fetch the area registry from Home Assistant, keyed by area ID."""
1163 # use_cache rebuilds the cached value from this return annotation, which rules out
1164 # the Area TypedDict; get_area_registry restores the type for callers
1165 return {area["area_id"]: area for area in await self.hass.get_area_registry()}
1166
1167
1168def _affects_mirrored_registry(data: Mapping[str, Any]) -> bool:
1169 """
1170 Return whether an entity registry update can change the mirrored entity registry.
1171
1172 :param data: The data of a Home Assistant entity_registry_updated event.
1173 """
1174 if data.get("action") != "update":
1175 # a created or removed entity always enters or leaves the listing
1176 return True
1177 # an update reports the fields it touched, so a change that only concerns fields we do
1178 # not mirror (a rename, an icon, a label) leaves our listing accurate. an update can also
1179 # report no fields at all, as a device rename re-derives a name field that Home Assistant
1180 # strips from the report, so treat that as a change of unknown reach
1181 if not (changes := data.get("changes")):
1182 return True
1183 return not REGISTRY_FIELDS_AFFECTING_MIRROR.isdisjoint(changes)
1184
1185
1186def _decompress_state(entity_id: str, compressed_state: CompressedState) -> State:
1187 """
1188 Return the full state representation of a compressed state message.
1189
1190 :param entity_id: The entity the compressed state belongs to.
1191 :param compressed_state: The compressed state as received over the websocket.
1192 """
1193 raw_context = compressed_state.get("c")
1194 context: Context = (
1195 raw_context
1196 if isinstance(raw_context, dict)
1197 else {"id": raw_context or "", "parent_id": None, "user_id": None}
1198 )
1199 last_changed = compressed_state.get("lc")
1200 # Home Assistant omits last_updated when it is identical to last_changed
1201 last_updated = compressed_state.get("lu", last_changed)
1202 return {
1203 "entity_id": entity_id,
1204 "state": compressed_state.get("s", ""),
1205 "attributes": compressed_state.get("a", {}),
1206 "last_changed": iso_from_utc_timestamp(last_changed) if last_changed else "",
1207 "last_updated": iso_from_utc_timestamp(last_updated) if last_updated else "",
1208 "context": context,
1209 }
1210