/
/
1"""Tests for the Home Assistant provider."""
2
3from __future__ import annotations
4
5import asyncio
6import json
7from collections.abc import AsyncIterator, Callable
8from contextlib import asynccontextmanager
9from math import ceil
10from typing import Any, cast
11from unittest.mock import AsyncMock, MagicMock, patch
12
13import pytest
14from aiohttp import ClientResponseError
15from hass_client.exceptions import BaseHassClientError
16from music_assistant_models.enums import EventType, ProviderFeature
17from music_assistant_models.errors import (
18 MusicAssistantError,
19 SetupFailedError,
20 UnsupportedFeaturedException,
21)
22
23from music_assistant.constants import CONF_LOG_LEVEL
24from music_assistant.providers.hass import (
25 CONF_AUTH_TOKEN,
26 CONF_URL,
27 CONF_VERIFY_SSL,
28 STATE_FETCH_BATCH_SIZE,
29 HassRegistryEntity,
30 HomeAssistantProvider,
31 setup,
32)
33from music_assistant.providers.hass.constants import (
34 CONF_MUTE_CONTROLS,
35 CONF_POWER_CONTROLS,
36 CONF_VOLUME_CONTROLS,
37)
38from tests.common import use_real_create_task
39
40LAST_CHANGED = 1683832716.072648
41LAST_CHANGED_ISO = "2023-05-11T19:18:36.072648+00:00"
42CONTEXT_ID = "01H0640ES8JCY1NGTNW3V41T5T"
43REGISTRY_LIST_COMMAND = "config/entity_registry/list_for_display"
44REGISTRY_ENTRIES_COMMAND = "config/entity_registry/get_entries"
45DEVICE_REGISTRY_LIST = "get_device_registry"
46
47
48def _state(entity_id: str, friendly_name: str) -> dict[str, Any]:
49 """Return a Home Assistant entity state."""
50 return {
51 "entity_id": entity_id,
52 "state": "idle",
53 "attributes": {"friendly_name": friendly_name},
54 }
55
56
57def _compressed(state: dict[str, Any]) -> dict[str, Any]:
58 """Return the compressed form Home Assistant sends for the given entity state."""
59 return {
60 "s": state["state"],
61 "a": state["attributes"],
62 "lc": LAST_CHANGED,
63 "c": CONTEXT_ID,
64 }
65
66
67def _config(**values: Any) -> MagicMock:
68 """Return a provider config exposing the given values via get_value (entry defaults)."""
69 persisted_values = {
70 CONF_URL: "http://homeassistant.local:8123",
71 CONF_AUTH_TOKEN: "token",
72 CONF_VERIFY_SSL: True,
73 CONF_LOG_LEVEL: "GLOBAL",
74 CONF_POWER_CONTROLS: [],
75 CONF_MUTE_CONTROLS: [],
76 CONF_VOLUME_CONTROLS: [],
77 **values,
78 }
79 config = MagicMock()
80 config.instance_id = "hass--test"
81 config.name = "Home Assistant"
82 config.get_value.side_effect = persisted_values.get
83 # get_setup_value falls through to config.values/get_value when setup_data is empty
84 config.values = {}
85 return config
86
87
88class _Cache:
89 """Provide the slice of the cache controller that @use_cache relies on."""
90
91 def __init__(self) -> None:
92 self.entries: dict[str, Any] = {}
93 self.fresh = True
94
95 async def get_with_freshness(self, key: str, **kwargs: Any) -> tuple[Any, bool, bool]:
96 """Return the (data, is_fresh, found) triplet for the given key."""
97 # the real controller reads the cache database here, so yield like it does:
98 # @use_cache stores in the background, and only a yield lets that store land
99 await asyncio.sleep(0)
100 if key not in self.entries:
101 return None, False, False
102 if not self.fresh and not kwargs.get("include_expired"):
103 return None, False, False
104 return self.entries[key], self.fresh, True
105
106 async def set(self, key: str, data: Any, **kwargs: Any) -> None:
107 """Store data under the given key."""
108 # the real controller serializes on a worker thread and then writes the cache
109 # database, so a store lands well after the call that scheduled it returned
110 await asyncio.sleep(0)
111 await asyncio.sleep(0)
112 self.entries[key] = data
113
114 def expire_all(self) -> None:
115 """Mark every stored entry as no longer fresh."""
116 self.fresh = False
117
118
119def _mass() -> MagicMock:
120 """Return the Music Assistant dependencies used during provider startup."""
121 mass = MagicMock()
122 mass.cache = _Cache()
123 mass.http_session = MagicMock()
124 mass.http_session_no_ssl = MagicMock()
125 use_real_create_task(mass)
126 mass.players.register_or_update_player_control = AsyncMock()
127 # get_setup_value reads the (empty, here) live setup_data blob from the store, then
128 # falls through to the provider config mock's get_value for the persisted test values
129 mass.config.get = MagicMock(return_value={})
130 mass.config.get_raw_provider_config_value = MagicMock(return_value=None)
131 return mass
132
133
134class _HomeAssistantClient:
135 """Provide lifecycle-aware Home Assistant behavior for provider tests."""
136
137 def __init__(
138 self,
139 states: list[dict[str, Any]],
140 registry_error: Exception | None = None,
141 listener_error: Exception | None = None,
142 connect_error: Exception | None = None,
143 block_registry: bool = False,
144 ) -> None:
145 self.connected = False
146 self.disconnected = False
147 self.listener_started = asyncio.Event()
148 self.listener_cancelled = asyncio.Event()
149 self.listener_stopped = asyncio.Event()
150 self.registry_started = asyncio.Event()
151 self.registry_cancelled = asyncio.Event()
152 self.registry_stopped = asyncio.Event()
153 self.calls: list[str] = []
154 self.subscribed = asyncio.Event()
155 # entity_ids per subscribe_entities call and per invoked unsubscribe callable
156 self.subscriptions: list[list[str]] = []
157 self.unsubscribed: list[list[str]] = []
158 self.active_subscriptions = 0
159 # (event_type, callback) per subscribe_events call
160 self.event_subscriptions: list[tuple[str, Callable[[dict[str, Any]], None]]] = []
161 self.active_event_subscriptions = 0
162 # compressed states keyed by entity_id, as delivered over the websocket
163 self.compressed_states = {state["entity_id"]: _compressed(state) for state in states}
164 # entities Home Assistant reports as disabled, so leaves out of the registry listing
165 self.disabled_entities: set[str] = set()
166 # devices as returned in full by the device registry listing
167 self.devices: list[dict[str, Any]] = []
168 # events delivered ahead of a subscription's initial state message
169 self.leading_events: list[dict[str, Any]] = []
170 self.deliver_initial_states = True
171 self._registry_error = registry_error
172 self._listener_error = listener_error
173 self._connect_error = connect_error
174 self._registry_result = (
175 asyncio.get_running_loop().create_future() if block_registry else None
176 )
177 # resolve this to make an already running listener return, as a lost connection does
178 self.connection_lost: asyncio.Future[None] = asyncio.get_running_loop().create_future()
179 self.send_command = AsyncMock(side_effect=self._send_command)
180
181 async def connect(self) -> None:
182 """Connect the client."""
183 self.calls.append("connect")
184 if self._connect_error:
185 raise self._connect_error
186 self.connected = True
187
188 async def start_listening(self) -> None:
189 """Listen until the connection is lost or the provider stops the listener task."""
190 self.calls.append("start_listening")
191 self.listener_started.set()
192 try:
193 if self._listener_error:
194 raise self._listener_error
195 await self.connection_lost
196 except asyncio.CancelledError:
197 self.listener_cancelled.set()
198 raise
199 finally:
200 if self._registry_result and not self._registry_result.done():
201 self._registry_result.cancel()
202 self.listener_stopped.set()
203
204 async def subscribe_entities(
205 self, cb_func: Callable[[dict[str, Any]], None], entity_ids: list[str]
206 ) -> Callable[[], None]:
207 """Deliver the subscription's state messages and return the unsubscribe callable."""
208 self.calls.append("subscribe_entities")
209 self.subscriptions.append(list(entity_ids))
210 self.active_subscriptions += 1
211 self.subscribed.set()
212 loop = asyncio.get_running_loop()
213 for event in self.leading_events:
214 loop.call_soon(cb_func, event)
215 if self.deliver_initial_states:
216 initial = {
217 entity_id: self.compressed_states[entity_id]
218 for entity_id in entity_ids
219 if entity_id in self.compressed_states
220 }
221 loop.call_soon(cb_func, {"a": initial})
222
223 def _unsubscribe() -> None:
224 self.calls.append("unsubscribe_entities")
225 self.unsubscribed.append(list(entity_ids))
226 self.active_subscriptions -= 1
227
228 return _unsubscribe
229
230 async def subscribe_events(
231 self, cb_func: Callable[[dict[str, Any]], None], event_type: str
232 ) -> Callable[[], None]:
233 """Register the event callback after command responses can be received."""
234 await self.listener_started.wait()
235 self.calls.append("subscribe_events")
236 self.event_subscriptions.append((event_type, cb_func))
237 self.active_event_subscriptions += 1
238
239 def _unsubscribe() -> None:
240 self.calls.append("unsubscribe_events")
241 self.active_event_subscriptions -= 1
242
243 return _unsubscribe
244
245 async def get_device_registry(self) -> list[dict[str, Any]]:
246 """Return the full device registry listing."""
247 self.calls.append(DEVICE_REGISTRY_LIST)
248 return list(self.devices)
249
250 def fire_event(self, event_type: str, data: dict[str, Any]) -> None:
251 """Deliver an event to every subscriber of the given event type."""
252 for subscribed_type, cb_func in self.event_subscriptions:
253 if subscribed_type == event_type:
254 cb_func({"event_type": event_type, "data": data})
255
256 async def disconnect(self) -> None:
257 """Disconnect the client."""
258 self.calls.append("disconnect")
259 self.connected = False
260 self.disconnected = True
261
262 async def _send_command(self, command: str, **kwargs: Any) -> Any:
263 """Return the response Home Assistant sends for the given websocket command."""
264 if command == REGISTRY_LIST_COMMAND:
265 return await self._registry_for_display()
266 self.calls.append(command)
267 if command == REGISTRY_ENTRIES_COMMAND:
268 return self._registry_entries(cast("list[str]", kwargs["entity_ids"]))
269 return {"response": {"data": "answer"}}
270
271 async def _registry_for_display(self) -> dict[str, Any]:
272 """Return the entity registry listing after command responses can be received."""
273 await self.listener_started.wait()
274 self.calls.append(REGISTRY_LIST_COMMAND)
275 self.registry_started.set()
276 try:
277 if self._registry_result:
278 await self._registry_result
279 if self._registry_error:
280 raise self._registry_error
281 return {
282 "entity_categories": {},
283 "entities": [
284 {"ei": entity_id, "pl": "test"}
285 for entity_id in self.compressed_states
286 if entity_id not in self.disabled_entities
287 ],
288 }
289 except asyncio.CancelledError:
290 self.registry_cancelled.set()
291 raise
292 finally:
293 self.registry_stopped.set()
294
295 def _registry_entries(self, entity_ids: list[str]) -> dict[str, dict[str, Any] | None]:
296 """Return the full registry entry of every requested entity, None when unknown."""
297 return {
298 entity_id: (
299 {
300 "entity_id": entity_id,
301 "id": f"registry_id_{entity_id}",
302 "platform": "test",
303 "device_id": "device_id",
304 "config_entry_id": "config_entry_id",
305 }
306 if entity_id in self.compressed_states
307 else None
308 )
309 for entity_id in entity_ids
310 }
311
312
313@asynccontextmanager
314async def _start_provider(
315 states: list[dict[str, Any]], **config_values: Any
316) -> AsyncIterator[tuple[HomeAssistantProvider, _HomeAssistantClient]]:
317 """Start the provider with a connected mocked Home Assistant client."""
318 hass = _HomeAssistantClient(states)
319 mass = _mass()
320 manifest = MagicMock()
321 manifest.domain = "hass"
322 manifest.name = "Home Assistant"
323 with patch("music_assistant.providers.hass.HomeAssistantClient", return_value=hass):
324 provider = await setup(mass, manifest, _config(**config_values))
325 assert isinstance(provider, HomeAssistantProvider)
326 async with asyncio.timeout(1):
327 await provider.handle_async_init()
328 try:
329 yield provider, hass
330 finally:
331 await provider.unload()
332
333
334async def _wait_for_stored(provider: HomeAssistantProvider) -> None:
335 """Wait until the background store of the device listing has landed in the cache."""
336 cache = cast("_Cache", provider.mass.cache)
337 async with asyncio.timeout(1):
338 while not cache.entries:
339 await asyncio.sleep(0)
340
341
342def _hold_back_registry_fetch(hass: _HomeAssistantClient) -> asyncio.Future[None]:
343 """Hold back the next registry listing and return the future that releases it."""
344 registry_response: asyncio.Future[None] = asyncio.get_running_loop().create_future()
345 hass._registry_result = registry_response
346 # the startup fetch left the event set, so re-arm it for the fetch under test
347 hass.registry_started.clear()
348 return registry_response
349
350
351async def _wait_for_registry_fetch(hass: _HomeAssistantClient) -> None:
352 """Wait until the held back registry listing is in flight."""
353 async with asyncio.timeout(1):
354 await hass.registry_started.wait()
355
356
357def _registry_event(
358 entity_id: str, action: str = "update", changes: dict[str, Any] | None = None
359) -> dict[str, Any]:
360 """Return the data Home Assistant sends in an entity_registry_updated event."""
361 data: dict[str, Any] = {"action": action, "entity_id": entity_id}
362 if action == "update":
363 # an update carries the old value of every field it touched
364 data["changes"] = changes or {}
365 return data
366
367
368async def _fire_registry_update(
369 provider: HomeAssistantProvider,
370 hass: _HomeAssistantClient,
371 entity_id: str,
372 action: str,
373) -> None:
374 """Deliver an entity registry update and wait for the engine rebuild it schedules."""
375 with patch("music_assistant.providers.hass.ENGINE_REFRESH_DEBOUNCE", 0):
376 hass.fire_event("entity_registry_updated", _registry_event(entity_id, action))
377 assert provider._engine_refresh_task is not None
378 async with asyncio.timeout(1):
379 await provider._engine_refresh_task
380
381
382def _providers_updated_events(provider: HomeAssistantProvider) -> list[Any]:
383 """
384 Return the PROVIDERS_UPDATED events the provider signalled so far.
385
386 :param provider: The provider whose Music Assistant mock is inspected.
387 """
388 signal_event = cast("MagicMock", provider.mass.signal_event)
389 return [
390 call
391 for call in signal_event.call_args_list
392 if call.args[:1] == (EventType.PROVIDERS_UPDATED,)
393 ]
394
395
396async def test_feature_resolution_starts_listener_first() -> None:
397 """Resolve startup features only after the Home Assistant listener starts."""
398 states = [
399 _state("ai_task.default", "Default AI"),
400 _state("tts.default", "Default TTS"),
401 ]
402
403 async with _start_provider(states) as (provider, hass):
404 assert hass.calls[:4] == [
405 "connect",
406 "start_listening",
407 "subscribe_events",
408 REGISTRY_LIST_COMMAND,
409 ]
410 assert ProviderFeature.AI_QUERY in provider.supported_features
411 assert ProviderFeature.TTS in provider.supported_features
412
413
414async def test_feature_resolution_failure_cleans_up_connection() -> None:
415 """Clean up the listener and connection when feature resolution fails."""
416 hass = _HomeAssistantClient([], BaseHassClientError("Unable to load Home Assistant states"))
417 manifest = MagicMock()
418 manifest.domain = "hass"
419 manifest.name = "Home Assistant"
420 with patch("music_assistant.providers.hass.HomeAssistantClient", return_value=hass):
421 provider = await setup(_mass(), manifest, _config())
422 assert isinstance(provider, HomeAssistantProvider)
423
424 with pytest.raises(SetupFailedError, match="Unable to load Home Assistant states"):
425 async with asyncio.timeout(1):
426 await provider.handle_async_init()
427
428 assert hass.listener_started.is_set()
429 assert hass.listener_stopped.is_set()
430 assert hass.disconnected
431 assert hass.calls == [
432 "connect",
433 "start_listening",
434 "subscribe_events",
435 REGISTRY_LIST_COMMAND,
436 "unsubscribe_events",
437 "disconnect",
438 ]
439 assert provider._listen_task is None
440
441
442async def test_listener_failure_does_not_mask_feature_resolution_failure() -> None:
443 """Preserve the startup error when the listener also fails."""
444 hass = _HomeAssistantClient(
445 [],
446 BaseHassClientError("Unable to load Home Assistant states"),
447 RuntimeError("Listener failed"),
448 )
449 manifest = MagicMock()
450 manifest.domain = "hass"
451 manifest.name = "Home Assistant"
452 with patch("music_assistant.providers.hass.HomeAssistantClient", return_value=hass):
453 provider = await setup(_mass(), manifest, _config())
454 assert isinstance(provider, HomeAssistantProvider)
455
456 with pytest.raises(SetupFailedError, match="Unable to load Home Assistant states"):
457 async with asyncio.timeout(1):
458 await provider.handle_async_init()
459
460 assert hass.disconnected
461 assert provider._listen_task is None
462
463
464async def test_listener_exit_terminates_pending_feature_resolution() -> None:
465 """Fail startup when the listener exits while feature resolution is pending."""
466 hass = _HomeAssistantClient(
467 [],
468 listener_error=BaseHassClientError("Listener failed"),
469 block_registry=True,
470 )
471 mass = _mass()
472 manifest = MagicMock()
473 manifest.domain = "hass"
474 manifest.name = "Home Assistant"
475 with patch("music_assistant.providers.hass.HomeAssistantClient", return_value=hass):
476 provider = await setup(mass, manifest, _config())
477 assert isinstance(provider, HomeAssistantProvider)
478
479 with pytest.raises(SetupFailedError, match="listener stopped during startup"):
480 async with asyncio.timeout(1):
481 await provider.handle_async_init()
482
483 assert hass.listener_started.is_set()
484 assert hass.listener_stopped.is_set()
485 assert hass.registry_stopped.is_set()
486 assert hass.disconnected
487 assert provider._listen_task is None
488 mass.call_later.assert_not_called()
489
490
491async def test_feature_resolution_timeout_cleans_up_connection() -> None:
492 """Clean up startup when Home Assistant feature resolution times out."""
493 hass = _HomeAssistantClient([], block_registry=True)
494 mass = _mass()
495 manifest = MagicMock()
496 manifest.domain = "hass"
497 manifest.name = "Home Assistant"
498 with (
499 patch("music_assistant.providers.hass.HomeAssistantClient", return_value=hass),
500 patch("music_assistant.providers.hass.FEATURE_DISCOVERY_TIMEOUT", 0.1),
501 ):
502 provider = await setup(mass, manifest, _config())
503 assert isinstance(provider, HomeAssistantProvider)
504 init_task = asyncio.create_task(provider.handle_async_init())
505 async with asyncio.timeout(1):
506 await hass.registry_started.wait()
507
508 with pytest.raises(
509 SetupFailedError, match="Timed out while resolving Home Assistant feature entities"
510 ):
511 await init_task
512
513 assert hass.registry_stopped.is_set()
514 assert hass.registry_cancelled.is_set()
515 assert hass.listener_stopped.is_set()
516 assert hass.listener_cancelled.is_set()
517 assert hass.disconnected
518 assert provider._listen_task is None
519 mass.call_later.assert_not_called()
520
521
522async def test_feature_resolution_cancellation_cleans_up_connection() -> None:
523 """Clean up startup when Home Assistant initialization is cancelled."""
524 hass = _HomeAssistantClient([], block_registry=True)
525 mass = _mass()
526 manifest = MagicMock()
527 manifest.domain = "hass"
528 manifest.name = "Home Assistant"
529 with patch("music_assistant.providers.hass.HomeAssistantClient", return_value=hass):
530 provider = await setup(mass, manifest, _config())
531 assert isinstance(provider, HomeAssistantProvider)
532 init_task = asyncio.create_task(provider.handle_async_init())
533 async with asyncio.timeout(1):
534 await hass.registry_started.wait()
535 init_task.cancel()
536
537 with pytest.raises(asyncio.CancelledError):
538 await init_task
539
540 assert hass.registry_stopped.is_set()
541 assert hass.registry_cancelled.is_set()
542 assert hass.listener_stopped.is_set()
543 assert hass.listener_cancelled.is_set()
544 assert hass.disconnected
545 assert provider._listen_task is None
546 mass.call_later.assert_not_called()
547
548
549async def test_connection_failure_cleans_up_client() -> None:
550 """Clean up the client when connecting to Home Assistant fails."""
551 hass = _HomeAssistantClient([], connect_error=BaseHassClientError("Unable to connect"))
552 manifest = MagicMock()
553 manifest.domain = "hass"
554 manifest.name = "Home Assistant"
555 with patch("music_assistant.providers.hass.HomeAssistantClient", return_value=hass):
556 provider = await setup(_mass(), manifest, _config())
557 assert isinstance(provider, HomeAssistantProvider)
558
559 with pytest.raises(SetupFailedError, match="Unable to connect"):
560 await provider.handle_async_init()
561
562 assert hass.disconnected
563 assert provider._listen_task is None
564
565
566async def test_lost_connection_arms_the_reconnect_under_the_load_task_id() -> None:
567 """A lost connection arms the reconnect so a (re)load starting first cancels it."""
568 async with _start_provider([]) as (provider, hass):
569 mass = cast("MagicMock", provider.mass)
570 mass.call_later.reset_mock()
571 assert provider._listen_task is not None
572
573 hass.connection_lost.set_exception(BaseHassClientError("connection lost"))
574 async with asyncio.timeout(1):
575 await provider._listen_task
576
577 assert provider.available is False
578 retry = mass.call_later.call_args
579 assert retry.args == (5, mass.load_provider, "hass--test")
580 assert retry.kwargs == {"allow_retry": True, "task_id": "load_provider_hass--test"}
581
582
583async def test_engines_are_listed_for_every_feature_entity() -> None:
584 """Expose every Home Assistant TTS and AI Task entity as an engine."""
585 states = [
586 _state("tts.piper", "Piper"),
587 _state("tts.cloud", "Home Assistant Cloud"),
588 _state("ai_task.openai", "OpenAI"),
589 _state("sensor.example", "Example"),
590 ]
591
592 async with _start_provider(states) as (provider, hass):
593 calls_before = list(hass.calls)
594 tts_engines = await provider.get_tts_engines()
595 ai_engines = await provider.get_ai_engines()
596
597 # listing engines is a hot path for consumers: it must not touch Home Assistant
598 assert hass.calls == calls_before
599 assert [(engine.id, engine.name) for engine in tts_engines] == [
600 ("tts.cloud", "Home Assistant Cloud (tts.cloud)"),
601 ("tts.piper", "Piper (tts.piper)"),
602 ]
603 assert [(engine.id, engine.name) for engine in ai_engines] == [
604 ("ai_task.openai", "OpenAI (ai_task.openai)")
605 ]
606 assert tts_engines[0].uid == "hass--test/tts.cloud"
607
608
609async def test_engine_without_friendly_name_falls_back_to_entity_id() -> None:
610 """Name an engine after its entity_id when Home Assistant has no friendly name."""
611 unnamed = {"entity_id": "tts.unnamed", "state": "idle", "attributes": {}}
612
613 async with _start_provider([unnamed]) as (provider, _):
614 engines = await provider.get_tts_engines()
615
616 assert [engine.name for engine in engines] == ["tts.unnamed"]
617
618
619async def test_ai_query_uses_the_first_engine_by_default() -> None:
620 """Send an AI query to the first available engine when none is requested."""
621 states = [_state("ai_task.first", "First"), _state("ai_task.second", "Second")]
622
623 async with _start_provider(states) as (provider, hass):
624 result = await provider.ai_query("What is this song?")
625
626 assert ProviderFeature.AI_QUERY in provider.supported_features
627 assert result == "answer"
628 hass.send_command.assert_awaited_with(
629 "call_service",
630 domain="ai_task",
631 service="generate_data",
632 service_data={
633 "task_name": "music_assistant",
634 "instructions": "What is this song?",
635 "entity_id": "ai_task.first",
636 },
637 return_response=True,
638 )
639
640
641async def test_ai_query_uses_the_requested_engine() -> None:
642 """Send an AI query to the requested engine."""
643 states = [_state("ai_task.first", "First"), _state("ai_task.second", "Second")]
644
645 async with _start_provider(states) as (provider, hass):
646 await provider.ai_query("What is this song?", engine_id="ai_task.second")
647
648 service_data = hass.send_command.call_args.kwargs["service_data"]
649 assert service_data["entity_id"] == "ai_task.second"
650
651
652async def test_ai_query_not_advertised_without_entity() -> None:
653 """Do not advertise AI queries when Home Assistant has no AI Task entity."""
654 async with _start_provider([_state("sensor.example", "Example")]) as (provider, _):
655 assert ProviderFeature.AI_QUERY not in provider.supported_features
656
657 with pytest.raises(UnsupportedFeaturedException):
658 await provider.ai_query("What is this song?")
659
660
661def _mock_tts_response(provider: HomeAssistantProvider) -> MagicMock:
662 """Let the Home Assistant tts_get_url endpoint return a URL and return the post mock."""
663 response = AsyncMock()
664 response.ok = True
665 response.raise_for_status = MagicMock()
666 response.json.return_value = {"url": "http://homeassistant.local/tts.mp3"}
667 post = cast("MagicMock", provider.mass.http_session.post)
668 post.return_value.__aenter__.return_value = response
669 return post
670
671
672def _mock_tts_error_response(provider: HomeAssistantProvider, error_message: str) -> MagicMock:
673 """Let the tts_get_url endpoint fail with a 400 carrying the given HA error body."""
674 response = AsyncMock()
675 response.ok = False
676 response.json.return_value = {"error": error_message}
677 response.raise_for_status = MagicMock(
678 side_effect=ClientResponseError(MagicMock(), (), status=400, message=error_message)
679 )
680 post = cast("MagicMock", provider.mass.http_session.post)
681 post.return_value.__aenter__.return_value = response
682 return post
683
684
685async def test_tts_uses_the_first_engine_by_default() -> None:
686 """Render speech on the first available engine when none is requested."""
687 states = [_state("tts.first", "First"), _state("tts.second", "Second")]
688
689 async with _start_provider(states) as (provider, _):
690 post = _mock_tts_response(provider)
691
692 stream = await provider.get_tts_message("Hello")
693
694 assert ProviderFeature.TTS in provider.supported_features
695 assert stream.path == "http://homeassistant.local/tts.mp3"
696 post.assert_called_once()
697 request = post.call_args
698 assert request.args == ("http://homeassistant.local:8123/api/tts_get_url",)
699 assert request.kwargs["json"] == {"engine_id": "tts.first", "message": "Hello"}
700
701
702async def test_tts_uses_the_requested_engine() -> None:
703 """Render speech on the requested engine."""
704 states = [_state("tts.first", "First"), _state("tts.second", "Second")]
705
706 async with _start_provider(states) as (provider, _):
707 post = _mock_tts_response(provider)
708
709 await provider.get_tts_message("Hello", engine_id="tts.second")
710
711 assert post.call_args.kwargs["json"] == {"engine_id": "tts.second", "message": "Hello"}
712
713
714async def test_tts_not_advertised_without_entity() -> None:
715 """Do not advertise TTS when Home Assistant has no TTS entity."""
716 async with _start_provider([_state("sensor.example", "Example")]) as (provider, _):
717 assert ProviderFeature.TTS not in provider.supported_features
718
719 with pytest.raises(UnsupportedFeaturedException):
720 await provider.get_tts_message("Hello")
721
722
723async def test_tts_sends_options_in_the_payload() -> None:
724 """A host's TTS options are forwarded in the tts_get_url payload."""
725 async with _start_provider([_state("tts.first", "First")]) as (provider, _):
726 post = _mock_tts_response(provider)
727
728 await provider.get_tts_message(
729 "Hello", options={"voice": "en_US-lessac-medium", "length_scale": 1.2}
730 )
731
732 assert post.call_args.kwargs["json"] == {
733 "engine_id": "tts.first",
734 "message": "Hello",
735 "options": {"voice": "en_US-lessac-medium", "length_scale": 1.2},
736 }
737
738
739async def test_tts_omits_options_from_the_payload_when_empty() -> None:
740 """An empty options dict is not sent to Home Assistant."""
741 async with _start_provider([_state("tts.first", "First")]) as (provider, _):
742 post = _mock_tts_response(provider)
743
744 await provider.get_tts_message("Hello", options={})
745
746 assert post.call_args.kwargs["json"] == {"engine_id": "tts.first", "message": "Hello"}
747
748
749async def test_tts_invalid_option_raises_music_assistant_error() -> None:
750 """HA rejecting an unknown TTS option surfaces as MusicAssistantError with HA's message."""
751 async with _start_provider([_state("tts.first", "First")]) as (provider, _):
752 _mock_tts_error_response(provider, "Invalid options found: ['speaking_cadance']")
753
754 with pytest.raises(MusicAssistantError, match=r"Invalid options found"):
755 await provider.get_tts_message("Hello", options={"speaking_cadance": 1})
756
757
758async def test_tts_error_body_that_is_not_json_still_raises_the_generic_way() -> None:
759 """An unparsable error body leaves the status error to speak for itself."""
760 async with _start_provider([_state("tts.first", "First")]) as (provider, _):
761 post = _mock_tts_error_response(provider, "Invalid options found: ['x']")
762 response = post.return_value.__aenter__.return_value
763 response.json.side_effect = ValueError("not json")
764
765 with pytest.raises(ClientResponseError):
766 await provider.get_tts_message("Hello", options={"x": 1})
767
768
769async def test_tts_unsupported_language_still_raises_the_generic_way() -> None:
770 """An unsupported-language 400 keeps raising as before, so the caller's retry still fires."""
771 async with _start_provider([_state("tts.first", "First")]) as (provider, _):
772 _mock_tts_error_response(provider, "Language 'xx' not supported")
773
774 with pytest.raises(ClientResponseError):
775 await provider.get_tts_message("Hello", language="xx")
776
777
778async def test_registry_update_refreshes_the_engines() -> None:
779 """Pick up a feature entity that Home Assistant adds after startup."""
780 async with _start_provider([_state("sensor.example", "Example")]) as (provider, hass):
781 await provider.loaded_in_mass()
782 assert ProviderFeature.TTS not in provider.supported_features
783
784 hass.compressed_states["tts.new"] = _compressed(_state("tts.new", "New"))
785 await _fire_registry_update(provider, hass, "tts.new", "create")
786
787 assert [engine.id for engine in await provider.get_tts_engines()] == ["tts.new"]
788 assert ProviderFeature.TTS in provider.supported_features
789
790
791async def test_registry_update_discards_the_feature_of_a_removed_engine() -> None:
792 """Stop advertising a feature once its last Home Assistant entity is gone."""
793 states = [_state("tts.only", "Only"), _state("ai_task.only", "Only")]
794
795 async with _start_provider(states) as (provider, hass):
796 await provider.loaded_in_mass()
797 assert ProviderFeature.TTS in provider.supported_features
798
799 del hass.compressed_states["tts.only"]
800 await _fire_registry_update(provider, hass, "tts.only", "remove")
801
802 assert await provider.get_tts_engines() == []
803 assert ProviderFeature.TTS not in provider.supported_features
804 # the AI Task entity is untouched, so its feature stays declared
805 assert ProviderFeature.AI_QUERY in provider.supported_features
806
807
808async def test_changed_engines_notify_the_consumers_once() -> None:
809 """Announce a refresh that changed the engine lists with a single PROVIDERS_UPDATED."""
810 states = [_state("tts.only", "Only"), _state("ai_task.only", "Only")]
811
812 async with _start_provider(states) as (provider, hass):
813 await provider.loaded_in_mass()
814 mass = cast("MagicMock", provider.mass)
815 mass.signal_event.reset_mock()
816
817 del hass.compressed_states["tts.only"]
818 await _fire_registry_update(provider, hass, "tts.only", "remove")
819
820 events = _providers_updated_events(provider)
821 assert len(events) == 1
822 assert events[0].kwargs["data"] is mass.get_providers.return_value
823
824
825async def test_a_lost_ai_engine_notifies_the_consumers() -> None:
826 """Announce a vanished AI engine, the selection AI Radio depends on."""
827 states = [_state("tts.only", "Only"), _state("ai_task.only", "Only")]
828
829 async with _start_provider(states) as (provider, hass):
830 await provider.loaded_in_mass()
831 cast("MagicMock", provider.mass).signal_event.reset_mock()
832
833 del hass.compressed_states["ai_task.only"]
834 await _fire_registry_update(provider, hass, "ai_task.only", "remove")
835
836 assert await provider.get_ai_engines() == []
837 assert [engine.id for engine in await provider.get_tts_engines()] == ["tts.only"]
838 assert len(_providers_updated_events(provider)) == 1
839
840
841async def test_engines_are_in_place_before_the_consumers_are_told() -> None:
842 """Expose the rebuilt engine lists before signalling, as consumers read them at once."""
843 states = [_state("tts.only", "Only"), _state("ai_task.only", "Only")]
844
845 async with _start_provider(states) as (provider, hass):
846 await provider.loaded_in_mass()
847 signal_event = cast("MagicMock", provider.mass.signal_event)
848 signal_event.reset_mock()
849 engines_when_told: list[list[str]] = []
850 signal_event.side_effect = lambda *_args, **_kwargs: engines_when_told.append(
851 [engine.id for engine in provider._tts_engines]
852 )
853
854 del hass.compressed_states["tts.only"]
855 await _fire_registry_update(provider, hass, "tts.only", "remove")
856
857 assert engines_when_told == [[]]
858
859
860async def test_unchanged_engines_do_not_notify_the_consumers() -> None:
861 """Stay silent when a refresh rebuilds the very same engine lists."""
862 states = [_state("tts.only", "Only"), _state("ai_task.only", "Only")]
863
864 async with _start_provider(states) as (provider, hass):
865 await provider.loaded_in_mass()
866 cast("MagicMock", provider.mass).signal_event.reset_mock()
867
868 # registry churn that leaves the feature entities as they are, as in a rename of
869 # an entity that the engine name does not depend on
870 await _fire_registry_update(provider, hass, "tts.only", "update")
871
872 assert [engine.id for engine in await provider.get_tts_engines()] == ["tts.only"]
873 assert [engine.id for engine in await provider.get_ai_engines()] == ["ai_task.only"]
874 assert _providers_updated_events(provider) == []
875
876
877async def test_startup_refresh_does_not_notify_the_consumers() -> None:
878 """Leave the announcement of the engines found during startup to the load path."""
879 states = [_state("tts.only", "Only"), _state("ai_task.only", "Only")]
880
881 async with _start_provider(states) as (provider, _):
882 # the refresh that filled the empty lists ran before startup was marked complete
883 assert provider._startup_complete
884 assert [engine.id for engine in await provider.get_tts_engines()] == ["tts.only"]
885 assert _providers_updated_events(provider) == []
886
887
888async def test_refresh_tracks_engines_and_features_in_both_directions() -> None:
889 """Follow the engine lists and their features as feature entities appear and vanish."""
890 async with _start_provider([_state("sensor.example", "Example")]) as (provider, hass):
891 await provider.loaded_in_mass()
892 cast("MagicMock", provider.mass).signal_event.reset_mock()
893 assert ProviderFeature.TTS not in provider.supported_features
894 assert ProviderFeature.AI_QUERY not in provider.supported_features
895
896 hass.compressed_states["tts.new"] = _compressed(_state("tts.new", "New TTS"))
897 hass.compressed_states["ai_task.new"] = _compressed(_state("ai_task.new", "New AI"))
898 await _fire_registry_update(provider, hass, "tts.new", "create")
899
900 assert [engine.id for engine in await provider.get_tts_engines()] == ["tts.new"]
901 assert [engine.id for engine in await provider.get_ai_engines()] == ["ai_task.new"]
902 assert ProviderFeature.TTS in provider.supported_features
903 assert ProviderFeature.AI_QUERY in provider.supported_features
904
905 del hass.compressed_states["tts.new"]
906 del hass.compressed_states["ai_task.new"]
907 await _fire_registry_update(provider, hass, "tts.new", "remove")
908
909 assert await provider.get_tts_engines() == []
910 assert await provider.get_ai_engines() == []
911 assert ProviderFeature.TTS not in provider.supported_features
912 assert ProviderFeature.AI_QUERY not in provider.supported_features
913 assert len(_providers_updated_events(provider)) == 2
914
915
916async def test_registry_update_of_another_domain_is_ignored() -> None:
917 """Ignore registry updates for entities that cannot back a feature."""
918 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
919 await provider.loaded_in_mass()
920
921 hass.fire_event("entity_registry_updated", _registry_event("light.kitchen", "create"))
922
923 assert provider._engine_refresh_task is None
924
925
926async def test_registry_update_of_a_feature_entity_refreshes_without_a_refetch() -> None:
927 """Rebuild the engine lists for a renamed feature entity without refetching the registry."""
928 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
929 await provider.loaded_in_mass()
930 registry = await provider.get_entity_registry()
931 hass.compressed_states["tts.only"] = _compressed(_state("tts.only", "Renamed"))
932
933 with patch("music_assistant.providers.hass.ENGINE_REFRESH_DEBOUNCE", 0):
934 hass.fire_event(
935 "entity_registry_updated", _registry_event("tts.only", changes={"name": "Only"})
936 )
937 assert provider._engine_refresh_task is not None
938 async with asyncio.timeout(1):
939 await provider._engine_refresh_task
940
941 engines = await provider.get_tts_engines()
942 assert [engine.name for engine in engines] == ["Renamed (tts.only)"]
943 assert provider._entity_registry is registry
944
945
946async def test_burst_of_registry_updates_triggers_a_single_refresh() -> None:
947 """Collect a burst of registry updates into one rebuild of the engine lists."""
948 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
949 await provider.loaded_in_mass()
950 registry_fetches = hass.calls.count(REGISTRY_LIST_COMMAND)
951
952 with patch("music_assistant.providers.hass.ENGINE_REFRESH_DEBOUNCE", 0.05):
953 for index in range(3):
954 hass.fire_event(
955 "entity_registry_updated", _registry_event(f"tts.new_{index}", "create")
956 )
957 assert provider._engine_refresh_task is not None
958 async with asyncio.timeout(1):
959 await provider._engine_refresh_task
960
961 assert hass.calls.count(REGISTRY_LIST_COMMAND) == registry_fetches + 1
962
963
964async def test_registry_update_of_another_domain_invalidates_the_registry() -> None:
965 """Refresh the cached registry for an entity that appeared in any domain."""
966 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
967 registry_fetches = hass.calls.count(REGISTRY_LIST_COMMAND)
968 hass.compressed_states["light.kitchen"] = _compressed(_state("light.kitchen", "Kitchen"))
969
970 hass.fire_event("entity_registry_updated", _registry_event("light.kitchen", "create"))
971
972 # an entity that can not back a feature must not trigger an engine rebuild
973 assert provider._engine_refresh_task is None
974 result = await provider.get_states(domains=("light",))
975 assert [state["entity_id"] for state in result] == ["light.kitchen"]
976 assert hass.calls.count(REGISTRY_LIST_COMMAND) == registry_fetches + 1
977
978
979@pytest.mark.parametrize(
980 "changes",
981 [
982 pytest.param({"name": "Old name"}, id="rename"),
983 pytest.param({"icon": None}, id="icon"),
984 pytest.param({"labels": []}, id="labels"),
985 pytest.param({"hidden_by": None}, id="hidden"),
986 # an integration reload re-registers its entities, which touches these
987 pytest.param({"capabilities": None, "supported_features": 0}, id="reload"),
988 ],
989)
990async def test_registry_update_of_unmirrored_fields_keeps_the_registry(
991 changes: dict[str, Any],
992) -> None:
993 """Keep the mirrored registry for an update that cannot change what it holds."""
994 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
995 registry = await provider.get_entity_registry()
996 registry_fetches = hass.calls.count(REGISTRY_LIST_COMMAND)
997
998 hass.fire_event(
999 "entity_registry_updated", _registry_event("light.kitchen", changes=changes)
1000 )
1001
1002 assert await provider.get_entity_registry() is registry
1003 assert hass.calls.count(REGISTRY_LIST_COMMAND) == registry_fetches
1004
1005
1006@pytest.mark.parametrize(
1007 "changes",
1008 [
1009 pytest.param({"entity_id": "light.old"}, id="entity_id"),
1010 pytest.param({"platform": "old_platform"}, id="platform"),
1011 pytest.param({"device_id": None}, id="device_id"),
1012 pytest.param({"area_id": None}, id="area_id"),
1013 # the listing omits disabled entities, so this one enters or leaves it
1014 pytest.param({"disabled_by": None}, id="disabled_by"),
1015 # a device rename reports no changed fields at all
1016 pytest.param({}, id="changes_empty"),
1017 # a move to another config entry can silently re-enable the entity
1018 pytest.param({"config_entry_id": "other"}, id="config_entry_id"),
1019 ],
1020)
1021async def test_registry_update_of_mirrored_fields_invalidates_the_registry(
1022 changes: dict[str, Any],
1023) -> None:
1024 """Refetch the mirrored registry for an update that can change what it holds."""
1025 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1026 await provider.get_entity_registry()
1027 registry_fetches = hass.calls.count(REGISTRY_LIST_COMMAND)
1028
1029 hass.fire_event(
1030 "entity_registry_updated", _registry_event("light.kitchen", changes=changes)
1031 )
1032
1033 assert provider._entity_registry is None
1034 await provider.get_entity_registry()
1035 assert hass.calls.count(REGISTRY_LIST_COMMAND) == registry_fetches + 1
1036
1037
1038async def test_unmirrored_change_during_the_fetch_is_still_cached() -> None:
1039 """Cache a registry listing that only an irrelevant registry update raced."""
1040 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1041 provider._entity_registry = None
1042 # hold back the registry response until the update has been delivered
1043 registry_response = _hold_back_registry_fetch(hass)
1044 lookup = asyncio.ensure_future(provider.get_states(domains=("tts",)))
1045 await _wait_for_registry_fetch(hass)
1046
1047 hass.fire_event(
1048 "entity_registry_updated", _registry_event("light.kitchen", changes={"name": "Old"})
1049 )
1050 registry_response.set_result(None)
1051 async with asyncio.timeout(1):
1052 await lookup
1053
1054 assert provider._entity_registry is not None
1055
1056
1057async def test_domains_are_resolved_through_the_display_registry() -> None:
1058 """Resolve domains through the compact registry listing only."""
1059 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1060 await provider.get_states(domains=("tts",))
1061
1062 assert REGISTRY_LIST_COMMAND in hass.calls
1063 assert "config/entity_registry/list" not in hass.calls
1064
1065
1066async def test_registry_is_fetched_once_per_connection() -> None:
1067 """Serve repeated domain lookups from a registry that is fetched only once."""
1068 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1069 registry_fetches = hass.calls.count(REGISTRY_LIST_COMMAND)
1070
1071 await provider.get_states(domains=("tts",))
1072 await provider.get_states(domains=("media_player",))
1073
1074 assert registry_fetches == 1
1075 assert hass.calls.count(REGISTRY_LIST_COMMAND) == registry_fetches
1076
1077
1078async def test_concurrent_domain_lookups_share_one_registry_fetch() -> None:
1079 """Let concurrent domain lookups share a single registry fetch."""
1080 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1081 registry_fetches = hass.calls.count(REGISTRY_LIST_COMMAND)
1082 provider._entity_registry = None
1083 # hold back the registry response until both lookups are waiting for it
1084 registry_response = _hold_back_registry_fetch(hass)
1085
1086 lookups = asyncio.gather(
1087 provider.get_states(domains=("tts",)), provider.get_states(domains=("tts",))
1088 )
1089 await _wait_for_registry_fetch(hass)
1090 registry_response.set_result(None)
1091 async with asyncio.timeout(1):
1092 results = await lookups
1093
1094 assert all(len(result) == 1 for result in results)
1095 assert hass.calls.count(REGISTRY_LIST_COMMAND) == registry_fetches + 1
1096
1097
1098async def test_registry_changed_during_the_fetch_is_not_cached() -> None:
1099 """Keep a registry listing out of the cache when a registry update outdated it in flight."""
1100 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1101 provider._entity_registry = None
1102 # hold back the registry response until the update has been delivered
1103 registry_response = _hold_back_registry_fetch(hass)
1104 lookup = asyncio.ensure_future(provider.get_states(domains=("tts",)))
1105 await _wait_for_registry_fetch(hass)
1106
1107 hass.fire_event("entity_registry_updated", _registry_event("light.kitchen", "create"))
1108 registry_response.set_result(None)
1109 async with asyncio.timeout(1):
1110 await lookup
1111
1112 assert provider._entity_registry is None
1113 registry_fetches = hass.calls.count(REGISTRY_LIST_COMMAND)
1114 await provider.get_states(domains=("tts",))
1115 assert hass.calls.count(REGISTRY_LIST_COMMAND) == registry_fetches + 1
1116
1117
1118async def test_shared_registry_rejects_writes() -> None:
1119 """Reject writes to the registry (and its entries) shared between all callers."""
1120 async with _start_provider([_state("tts.only", "Only")]) as (provider, _):
1121 registry = await provider.get_entity_registry()
1122
1123 with pytest.raises(TypeError):
1124 registry["light.kitchen"] = HassRegistryEntity( # type: ignore[index]
1125 platform="test", device_id=None, area_id=None
1126 )
1127 with pytest.raises(AttributeError):
1128 registry["tts.only"].platform = "test" # type: ignore[misc]
1129
1130
1131async def test_registry_reuses_repeated_strings() -> None:
1132 """Hold on to a single string object per distinct platform, device id and area id."""
1133 async with _start_provider([_state("tts.only", "Only")]) as (provider, _):
1134 # decode the listing like a real response, so the repeated platform, device id and
1135 # area id arrive as distinct string objects instead of shared literals
1136 entities = json.loads(
1137 json.dumps(
1138 [
1139 {
1140 "ei": f"light.lamp_{index}",
1141 "pl": "esphome",
1142 "di": "device",
1143 "ai": "area",
1144 }
1145 for index in range(3)
1146 ]
1147 )
1148 )
1149 with patch.object(
1150 provider.hass, "send_command", AsyncMock(return_value={"entities": entities})
1151 ):
1152 registry = await provider._fetch_entity_registry()
1153
1154 assert len(registry) == 3
1155 assert registry["light.lamp_0"] == HassRegistryEntity(
1156 platform="esphome", device_id="device", area_id="area"
1157 )
1158 assert len({id(entry.platform) for entry in registry.values()}) == 1
1159 assert len({id(entry.device_id) for entry in registry.values()}) == 1
1160 assert len({id(entry.area_id) for entry in registry.values()}) == 1
1161
1162
1163async def test_registry_leaves_out_the_device_and_area_it_is_not_told_about() -> None:
1164 """Report no device or area for the entities whose listing entry omits them."""
1165 async with _start_provider([_state("tts.only", "Only")]) as (provider, _):
1166 entities = [
1167 {"ei": "light.full", "pl": "esphome", "di": "device", "ai": "area"},
1168 {"ei": "light.bare", "pl": "esphome"},
1169 {"ei": "light.device_only", "pl": "esphome", "di": "other_device"},
1170 ]
1171 with patch.object(
1172 provider.hass, "send_command", AsyncMock(return_value={"entities": entities})
1173 ):
1174 registry = await provider._fetch_entity_registry()
1175
1176 assert registry["light.bare"] == HassRegistryEntity(
1177 platform="esphome", device_id=None, area_id=None
1178 )
1179 assert registry["light.device_only"] == HassRegistryEntity(
1180 platform="esphome", device_id="other_device", area_id=None
1181 )
1182
1183
1184async def test_device_registry_is_reused_within_the_cache_window() -> None:
1185 """Serve a later device lookup from the cached listing."""
1186 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1187 hass.devices = [{"id": "dev1"}]
1188
1189 assert await provider.get_device_registry() == {"dev1": {"id": "dev1"}}
1190 await _wait_for_stored(provider)
1191 await provider.get_device_registry()
1192
1193 assert hass.calls.count(DEVICE_REGISTRY_LIST) == 1
1194
1195
1196async def test_device_registry_is_refetched_once_the_cache_window_passed() -> None:
1197 """Fetch the device listing again once the cached entry is no longer fresh."""
1198 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1199 hass.devices = [{"id": "dev1"}]
1200 await provider.get_device_registry()
1201 hass.devices = [{"id": "dev1"}, {"id": "dev2"}]
1202
1203 cast("_Cache", provider.mass.cache).expire_all()
1204
1205 assert list(await provider.get_device_registry()) == ["dev1", "dev2"]
1206 assert hass.calls.count(DEVICE_REGISTRY_LIST) == 2
1207
1208
1209async def test_device_entries_survive_the_cache_round_trip() -> None:
1210 """Return a cached device entry with all of its fields intact."""
1211 device = {
1212 "id": "dev1",
1213 "name": "Kitchen Speaker",
1214 "name_by_user": None,
1215 "connections": [["mac", "aa:bb:cc:dd:ee:ff"]],
1216 "manufacturer": "ESPHome",
1217 }
1218 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1219 hass.devices = [device]
1220
1221 fetched = await provider.get_device_registry()
1222 await _wait_for_stored(provider)
1223 cached = await provider.get_device_registry()
1224
1225 assert hass.calls.count(DEVICE_REGISTRY_LIST) == 1
1226 # the cached read is reconstructed from the return annotation, so it must not
1227 # drop or reshape any of the device fields
1228 assert cached == fetched == {"dev1": device}
1229
1230
1231async def test_concurrent_device_lookups_do_not_fetch_per_caller() -> None:
1232 """Keep a burst of concurrent device lookups off a fetch-per-caller path."""
1233 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1234 hass.devices = [{"id": "dev1"}]
1235
1236 async with asyncio.timeout(1):
1237 results = await asyncio.gather(*(provider.get_device_registry() for _ in range(20)))
1238
1239 assert all(list(result) == ["dev1"] for result in results)
1240 # the lock keeps the fetch count flat instead of growing with the burst; it does
1241 # not reach one, because use_cache stores in the background and the caller right
1242 # behind the first one still finds the cache cold
1243 assert hass.calls.count(DEVICE_REGISTRY_LIST) <= 2
1244
1245
1246async def test_disabled_entity_is_never_requested() -> None:
1247 """Leave an entity that Home Assistant disabled out of the state fetch."""
1248 states = [_state("media_player.kitchen", "Kitchen"), _state("media_player.spare", "Spare")]
1249
1250 async with _start_provider(states) as (provider, hass):
1251 # Home Assistant omits disabled entities from the registry, their state remains
1252 hass.disabled_entities.add("media_player.spare")
1253 hass.fire_event(
1254 "entity_registry_updated",
1255 _registry_event("media_player.spare", changes={"disabled_by": None}),
1256 )
1257 hass.subscriptions.clear()
1258
1259 result = await provider.get_states(domains=("media_player",))
1260
1261 assert [state["entity_id"] for state in result] == ["media_player.kitchen"]
1262 assert hass.subscriptions == [["media_player.kitchen"]]
1263
1264
1265async def test_registry_entries_are_fetched_for_the_given_entities_only() -> None:
1266 """Fetch the full registry entries of exactly the requested entities."""
1267 async with _start_provider([_state("media_player.kitchen", "Kitchen")]) as (provider, hass):
1268 entries = await provider.get_entity_registry_entries(
1269 ["media_player.kitchen", "media_player.unknown"]
1270 )
1271
1272 # an entity Home Assistant does not know is absent from the result
1273 assert list(entries) == ["media_player.kitchen"]
1274 assert entries["media_player.kitchen"]["config_entry_id"] == "config_entry_id"
1275 hass.send_command.assert_awaited_with(
1276 REGISTRY_ENTRIES_COMMAND,
1277 entity_ids=["media_player.kitchen", "media_player.unknown"],
1278 )
1279
1280
1281async def test_registry_entries_of_nothing_skips_the_round_trip() -> None:
1282 """Do not contact Home Assistant when no entity is requested."""
1283 async with _start_provider([_state("media_player.kitchen", "Kitchen")]) as (provider, hass):
1284 calls_before = list(hass.calls)
1285
1286 assert await provider.get_entity_registry_entries([]) == {}
1287 assert hass.calls == calls_before
1288
1289
1290async def test_entity_registry_subscription_is_replaced() -> None:
1291 """Replace the entity registry subscription instead of stacking a second one."""
1292 async with _start_provider([_state("tts.only", "Only")]) as (provider, hass):
1293 assert hass.active_event_subscriptions == 1
1294 assert hass.event_subscriptions[0][0] == "entity_registry_updated"
1295
1296 await provider._subscribe_entity_registry()
1297
1298 assert hass.active_event_subscriptions == 1
1299 assert hass.calls[-2:] == ["unsubscribe_events", "subscribe_events"]
1300
1301 await provider.unload()
1302
1303 assert hass.active_event_subscriptions == 0
1304
1305
1306async def test_config_entries_do_not_read_home_assistant() -> None:
1307 """Build the config entries without asking Home Assistant for anything."""
1308 async with _start_provider([_state("switch.example", "Example")]) as (provider, _):
1309 provider.available = True
1310 with patch.object(provider, "get_states", AsyncMock()) as get_states:
1311 entries = await provider.get_config_entries()
1312
1313 get_states.assert_not_awaited()
1314 assert CONF_POWER_CONTROLS in {entry.key for entry in entries}
1315
1316
1317async def test_config_entries_offer_the_control_lists_without_options() -> None:
1318 """Offer the control lists as plain entity id lists, filled in by the entity picker."""
1319 async with _start_provider([_state("switch.example", "Example")]) as (provider, _):
1320 provider.available = True
1321 entries = await provider.get_config_entries()
1322
1323 control_entries = [
1324 entry
1325 for entry in entries
1326 if entry.key in (CONF_POWER_CONTROLS, CONF_VOLUME_CONTROLS, CONF_MUTE_CONTROLS)
1327 ]
1328 assert len(control_entries) == 3
1329 for entry in control_entries:
1330 assert entry.options == []
1331 assert entry.multi_value is True
1332
1333
1334async def test_domain_states_use_a_single_subscription() -> None:
1335 """Fetch every entity of a domain in one websocket round-trip."""
1336 states = [_state(f"media_player.player_{index}", f"Player {index}") for index in range(50)]
1337
1338 async with _start_provider(states) as (provider, hass):
1339 result = await provider.get_states(domains=("media_player",))
1340
1341 assert len(result) == len(states)
1342 assert len(hass.subscriptions) == 1
1343 assert hass.subscriptions[0] == sorted(state["entity_id"] for state in states)
1344
1345
1346async def test_large_requests_are_split_into_batches() -> None:
1347 """Split a request that exceeds the batch size into bounded batches."""
1348 entity_ids = [
1349 f"media_player.player_{index:04d}" for index in range(STATE_FETCH_BATCH_SIZE * 2 + 1)
1350 ]
1351 states = [_state(entity_id, entity_id) for entity_id in entity_ids]
1352
1353 async with _start_provider(states) as (provider, hass):
1354 result = await provider.get_states(domains=("media_player",))
1355
1356 assert len(result) == len(entity_ids)
1357 assert len(hass.subscriptions) == ceil(len(entity_ids) / STATE_FETCH_BATCH_SIZE)
1358 assert all(len(batch) <= STATE_FETCH_BATCH_SIZE for batch in hass.subscriptions)
1359 requested = [entity_id for batch in hass.subscriptions for entity_id in batch]
1360 assert sorted(requested) == sorted(entity_ids)
1361 assert len(requested) == len(set(requested))
1362
1363
1364async def test_every_batch_is_unsubscribed() -> None:
1365 """Release the subscription of every batch once its states have been received."""
1366 entity_ids = [f"media_player.player_{index}" for index in range(5)]
1367 states = [_state(entity_id, entity_id) for entity_id in entity_ids]
1368
1369 async with _start_provider(states) as (provider, hass):
1370 with patch("music_assistant.providers.hass.STATE_FETCH_BATCH_SIZE", 2):
1371 await provider.get_states(domains=("media_player",))
1372
1373 assert len(hass.subscriptions) == 3
1374 assert hass.unsubscribed == hass.subscriptions
1375 assert hass.active_subscriptions == 0
1376
1377
1378async def test_timed_out_fetch_is_unsubscribed() -> None:
1379 """Release the subscription when Home Assistant never sends the states."""
1380 async with _start_provider([_state("media_player.kitchen", "Kitchen")]) as (provider, hass):
1381 hass.deliver_initial_states = False
1382
1383 with (
1384 patch("music_assistant.providers.hass.STATE_FETCH_TIMEOUT", 0.05),
1385 pytest.raises(TimeoutError),
1386 ):
1387 await provider.get_states(entity_ids=["media_player.kitchen"])
1388
1389 assert hass.unsubscribed == [["media_player.kitchen"]]
1390 assert hass.active_subscriptions == 0
1391
1392
1393async def test_cancelled_fetch_is_unsubscribed() -> None:
1394 """Release the subscription when the state fetch is cancelled."""
1395 async with _start_provider([_state("media_player.kitchen", "Kitchen")]) as (provider, hass):
1396 hass.deliver_initial_states = False
1397 fetch_task = asyncio.create_task(provider.get_states(entity_ids=["media_player.kitchen"]))
1398 async with asyncio.timeout(1):
1399 await hass.subscribed.wait()
1400 fetch_task.cancel()
1401
1402 with pytest.raises(asyncio.CancelledError):
1403 await fetch_task
1404
1405 assert hass.unsubscribed == [["media_player.kitchen"]]
1406 assert hass.active_subscriptions == 0
1407
1408
1409async def test_compressed_states_are_expanded() -> None:
1410 """Expand the compressed states of the initial message into full states."""
1411 async with _start_provider([]) as (provider, hass):
1412 hass.compressed_states = {
1413 "media_player.full": {
1414 "s": "playing",
1415 "a": {"friendly_name": "Full"},
1416 "lc": LAST_CHANGED,
1417 "lu": 1683838800.736819,
1418 "c": {"id": CONTEXT_ID, "parent_id": None, "user_id": "user"},
1419 },
1420 "media_player.unchanged": {"s": "idle", "lc": LAST_CHANGED, "c": CONTEXT_ID},
1421 "media_player.minimal": {},
1422 }
1423
1424 result = {
1425 state["entity_id"]: state
1426 for state in await provider.get_states(entity_ids=list(hass.compressed_states))
1427 }
1428
1429 assert result["media_player.full"]["state"] == "playing"
1430 assert result["media_player.full"]["attributes"] == {"friendly_name": "Full"}
1431 assert result["media_player.full"]["last_changed"] == LAST_CHANGED_ISO
1432 assert result["media_player.full"]["last_updated"] == "2023-05-11T21:00:00.736819+00:00"
1433 assert result["media_player.full"]["context"] == {
1434 "id": CONTEXT_ID,
1435 "parent_id": None,
1436 "user_id": "user",
1437 }
1438 # last_updated is omitted by HA when it is identical to last_changed
1439 assert result["media_player.unchanged"]["last_changed"] == LAST_CHANGED_ISO
1440 assert result["media_player.unchanged"]["last_updated"] == LAST_CHANGED_ISO
1441 assert result["media_player.unchanged"]["context"] == {
1442 "id": CONTEXT_ID,
1443 "parent_id": None,
1444 "user_id": None,
1445 }
1446 assert result["media_player.minimal"] == {
1447 "entity_id": "media_player.minimal",
1448 "state": "",
1449 "attributes": {},
1450 "last_changed": "",
1451 "last_updated": "",
1452 "context": {"id": "", "parent_id": None, "user_id": None},
1453 }
1454
1455
1456async def test_entity_without_state_is_absent() -> None:
1457 """Omit entities that Home Assistant has no state for."""
1458 async with _start_provider([_state("media_player.kitchen", "Kitchen")]) as (provider, hass):
1459 result = await provider.get_states(
1460 entity_ids=["media_player.kitchen", "media_player.removed"]
1461 )
1462
1463 assert [state["entity_id"] for state in result] == ["media_player.kitchen"]
1464 assert hass.subscriptions == [["media_player.kitchen", "media_player.removed"]]
1465
1466
1467async def test_state_change_does_not_complete_the_fetch() -> None:
1468 """Ignore a state change that arrives before the initial state message."""
1469 async with _start_provider([_state("media_player.kitchen", "Kitchen")]) as (provider, hass):
1470 hass.leading_events = [{"c": {"media_player.kitchen": {"+": {"s": "playing"}}}}]
1471
1472 result = await provider.get_states(entity_ids=["media_player.kitchen"])
1473
1474 assert [state["entity_id"] for state in result] == ["media_player.kitchen"]
1475 assert result[0]["state"] == "idle"
1476
1477
1478async def test_player_control_subscription_is_replaced() -> None:
1479 """Replace the player control subscription instead of stacking a second one."""
1480 states = [_state("media_player.kitchen", "Kitchen")]
1481 async with _start_provider(states, **{CONF_POWER_CONTROLS: ["media_player.kitchen"]}) as (
1482 provider,
1483 hass,
1484 ):
1485 await provider.loaded_in_mass()
1486 assert hass.active_subscriptions == 1
1487
1488 # only a changed selection results in a new subscription
1489 provider.config = _config(
1490 **{
1491 CONF_POWER_CONTROLS: ["media_player.kitchen"],
1492 CONF_MUTE_CONTROLS: ["switch.kitchen_amp"],
1493 }
1494 )
1495 await provider._register_player_controls()
1496
1497 assert len(hass.subscriptions) == 4
1498 assert hass.active_subscriptions == 1
1499 # the previous control subscription is only released once its replacement is live
1500 assert hass.calls[-2:] == ["subscribe_entities", "unsubscribe_entities"]
1501
1502 await provider.unload()
1503
1504 assert hass.active_subscriptions == 0
1505
1506
1507async def test_failed_control_subscription_keeps_the_previous_one() -> None:
1508 """A failed subscription attempt leaves the controls watched by the earlier one."""
1509 states = [_state("media_player.kitchen", "Kitchen"), _state("switch.amp", "Amp")]
1510 async with _start_provider(states, **{CONF_POWER_CONTROLS: ["media_player.kitchen"]}) as (
1511 provider,
1512 hass,
1513 ):
1514 await provider.loaded_in_mass()
1515 unsubscribe_before = provider._unsubscribe_controls
1516 assert hass.active_subscriptions == 1
1517
1518 async def _failing_subscribe(*_args: Any, **_kwargs: Any) -> Callable[[], None]:
1519 raise BaseHassClientError("connection lost")
1520
1521 hass.subscribe_entities = _failing_subscribe # type: ignore[method-assign]
1522 provider.config = _config(
1523 **{
1524 CONF_POWER_CONTROLS: ["media_player.kitchen"],
1525 CONF_MUTE_CONTROLS: ["switch.amp"],
1526 }
1527 )
1528
1529 with pytest.raises(BaseHassClientError):
1530 await provider._register_player_controls()
1531
1532 assert provider._unsubscribe_controls is unsubscribe_before
1533 assert hass.active_subscriptions == 1
1534
1535
1536async def test_unchanged_control_selection_skips_home_assistant() -> None:
1537 """Reconciling an unchanged selection does not talk to Home Assistant at all."""
1538 states = [_state("media_player.kitchen", "Kitchen")]
1539 async with _start_provider(states, **{CONF_POWER_CONTROLS: ["media_player.kitchen"]}) as (
1540 provider,
1541 hass,
1542 ):
1543 await provider.loaded_in_mass()
1544 calls_after_load = list(hass.calls)
1545
1546 await provider._register_player_controls()
1547
1548 assert hass.calls == calls_after_load
1549
1550
1551async def test_control_list_change_reconciles_without_reload() -> None:
1552 """A change confined to the control lists is applied without reloading the provider."""
1553 states = [_state("media_player.kitchen", "Kitchen"), _state("switch.amp", "Amp")]
1554 async with _start_provider(states, **{CONF_POWER_CONTROLS: ["media_player.kitchen"]}) as (
1555 provider,
1556 _hass,
1557 ):
1558 mass = cast("MagicMock", provider.mass)
1559 await provider.loaded_in_mass()
1560 assert set(provider._player_controls or {}) == {"media_player.kitchen"}
1561 mass.call_later.reset_mock()
1562
1563 await provider.update_config(
1564 _config(**{CONF_POWER_CONTROLS: ["switch.amp"]}),
1565 {f"values/{CONF_POWER_CONTROLS}"},
1566 )
1567
1568 mass.call_later.assert_not_called()
1569 assert set(provider._player_controls or {}) == {"switch.amp"}
1570 mass.players.remove_player_control.assert_called_once_with("media_player.kitchen")
1571 registered = mass.players.register_or_update_player_control.call_args[0][0]
1572 assert registered.id == "switch.amp"
1573 assert registered.supports_power is True
1574
1575
1576async def test_control_role_change_is_applied() -> None:
1577 """Moving an already selected entity to another role re-registers it in that role."""
1578 states = [_state("media_player.kitchen", "Kitchen")]
1579 async with _start_provider(states, **{CONF_POWER_CONTROLS: ["media_player.kitchen"]}) as (
1580 provider,
1581 _hass,
1582 ):
1583 await provider.loaded_in_mass()
1584
1585 await provider.update_config(
1586 _config(**{CONF_VOLUME_CONTROLS: ["media_player.kitchen"]}),
1587 {f"values/{CONF_POWER_CONTROLS}", f"values/{CONF_VOLUME_CONTROLS}"},
1588 )
1589
1590 control = (provider._player_controls or {})["media_player.kitchen"]
1591 assert control.supports_volume is True
1592 assert control.supports_power is False
1593
1594
1595async def test_control_selection_drops_a_value_that_is_no_entity_id() -> None:
1596 """A stored selection that is no entity ID must not take the other controls down."""
1597 states = [_state("media_player.kitchen", "Kitchen")]
1598 async with _start_provider(
1599 states, **{CONF_POWER_CONTROLS: ["power_controls", "media_player.kitchen"]}
1600 ) as (provider, hass):
1601 await provider.loaded_in_mass()
1602
1603 assert set(provider._player_controls or {}) == {"media_player.kitchen"}
1604 # Home Assistant refuses the whole subscription when it carries the stray value
1605 assert hass.subscriptions[-1] == ["media_player.kitchen"]
1606
1607
1608@pytest.mark.parametrize(
1609 "changed_keys",
1610 [
1611 {f"values/{CONF_URL}"},
1612 # a control list changing alongside another value still needs the reload
1613 {f"values/{CONF_POWER_CONTROLS}", f"values/{CONF_URL}"},
1614 ],
1615)
1616async def test_other_config_change_reloads_provider(changed_keys: set[str]) -> None:
1617 """Any change beyond the control lists reloads the provider."""
1618 async with _start_provider([]) as (provider, _hass):
1619 mass = cast("MagicMock", provider.mass)
1620 await provider.loaded_in_mass()
1621 mass.call_later.reset_mock()
1622
1623 await provider.update_config(_config(**{CONF_POWER_CONTROLS: ["switch.amp"]}), changed_keys)
1624
1625 mass.call_later.assert_called_once()
1626 assert mass.call_later.call_args[0][1] is mass.load_provider_config
1627