diff --git a/homeassistant/components/media_player/llm.py b/homeassistant/components/media_player/llm.py index 5c4114472119..d11e680d44c2 100644 --- a/homeassistant/components/media_player/llm.py +++ b/homeassistant/components/media_player/llm.py @@ -1,46 +1,63 @@ """LLM tools for the media_player integration.""" +from typing import Any, cast, override + +import probatio + from homeassistant.components.homeassistant import async_should_expose from homeassistant.components.llm import LLMTools -from homeassistant.core import HomeAssistant, callback -from homeassistant.helpers import intent +from homeassistant.const import ATTR_ENTITY_ID, ATTR_SUPPORTED_FEATURES +from homeassistant.core import HomeAssistant, State, callback +from homeassistant.helpers import config_validation as cv, intent from homeassistant.helpers.llm import ( LLM_API_ASSIST, IntentTool, LLMContext, Tool, ToolAnnotations, + ToolInput, + ToolResult, + async_get_match_preferences, ) +from homeassistant.util.json import JsonValueType +from .browse_media import SearchMedia from .const import ( + ATTR_MEDIA_CONTENT_ID, + ATTR_MEDIA_CONTENT_TYPE, + ATTR_MEDIA_FILTER_CLASSES, + ATTR_MEDIA_SEARCH_QUERY, DOMAIN, INTENT_MEDIA_NEXT, INTENT_MEDIA_PAUSE, INTENT_MEDIA_PREVIOUS, - INTENT_MEDIA_SEARCH_AND_PLAY, INTENT_MEDIA_UNPAUSE, INTENT_PLAYER_MUTE, INTENT_PLAYER_UNMUTE, INTENT_SET_VOLUME, INTENT_SET_VOLUME_RELATIVE, + SERVICE_PLAY_MEDIA, + SERVICE_SEARCH_MEDIA, + MediaClass, + MediaPlayerEntityFeature, ) # Intents owned by this integration that are exposed as LLM tools. +# HassMediaSearchAndPlay is not listed because the search media and play media +# tools cover it. LLM_INTENTS = { INTENT_MEDIA_NEXT: "Next track", INTENT_MEDIA_PAUSE: "Pause media", INTENT_PLAYER_MUTE: "Mute player", INTENT_PLAYER_UNMUTE: "Unmute player", INTENT_MEDIA_PREVIOUS: "Previous track", - INTENT_MEDIA_SEARCH_AND_PLAY: "Search and play media", INTENT_MEDIA_UNPAUSE: "Resume media", INTENT_SET_VOLUME: "Set volume", INTENT_SET_VOLUME_RELATIVE: "Change volume", } # Setting a value on the user's own player has no further effect when it is -# repeated. Stepping through tracks or volume has an effect on every call, and -# a search reaches the media the player can read. +# repeated. Stepping through tracks or volume has an effect on every call. _CONTROL = ToolAnnotations(idempotent=True, open_world=False) _CUMULATIVE = ToolAnnotations(open_world=False) @@ -54,25 +71,198 @@ INTENT_ANNOTATIONS = { INTENT_MEDIA_NEXT: _CUMULATIVE, INTENT_MEDIA_PREVIOUS: _CUMULATIVE, INTENT_SET_VOLUME_RELATIVE: _CUMULATIVE, - INTENT_MEDIA_SEARCH_AND_PLAY: ToolAnnotations(), } +# Both tools match players with the same features, so the same target +# arguments resolve to the same player. Search results are only valid on the +# player that returned them. +SEARCH_PLAY_FEATURES = ( + MediaPlayerEntityFeature.SEARCH_MEDIA | MediaPlayerEntityFeature.PLAY_MEDIA +) + +# Some players return hundreds of results, which would fill the LLM context. +MAX_SEARCH_RESULTS = 20 + +TARGET_SCHEMA = { + probatio.Optional("player_name"): cv.string, + probatio.Optional("player_area"): cv.string, + probatio.Optional("player_floor"): cv.string, +} + + +def _validate_args( + schema: probatio.Schema, tool_args: dict[str, Any] +) -> dict[str, Any]: + """Validate tool arguments, omitting the blank values that LLMs often send.""" + args: dict[str, Any] = schema( + { + key: value + for key, value in tool_args.items() + if not intent.is_blank_slot_value(value) + } + ) + return args + + +@callback +def _async_match_player( + hass: HomeAssistant, llm_context: LLMContext, args: dict[str, Any] +) -> State: + """Return the single media player that the target arguments match.""" + constraints = intent.MatchTargetsConstraints( + name=args.get("player_name"), + area_name=args.get("player_area"), + floor_name=args.get("player_floor"), + domains={DOMAIN}, + assistant=llm_context.assistant, + features=SEARCH_PLAY_FEATURES, + single_target=True, + ) + preferences = async_get_match_preferences(hass, llm_context) + result = intent.async_match_targets(hass, constraints, preferences) + if not result.is_match: + raise intent.MatchFailedError( + result=result, constraints=constraints, preferences=preferences + ) + return result.states[0] + + +class MediaSearchTool(Tool): + """LLM Tool that searches a media player for media.""" + + name = "media_player__search_media" + title = "Search media" + description = "Searches a media player for media and returns the playable items." + parameters = probatio.Schema( + { + probatio.Required( + ATTR_MEDIA_SEARCH_QUERY, + description="What to search for, such as a song, artist or album", + ): cv.string, + probatio.Optional( + "media_class", description="Only return media of this class" + ): probatio.In([cls.value for cls in MediaClass]), + **TARGET_SCHEMA, + } + ) + annotations = ToolAnnotations(read_only=True, destructive=False, idempotent=True) + integration = DOMAIN + + @override + async def async_call( + self, hass: HomeAssistant, tool_input: ToolInput, llm_context: LLMContext + ) -> ToolResult: + """Search a media player.""" + args = _validate_args(self.parameters, tool_input.tool_args) + entity_id = _async_match_player(hass, llm_context, args).entity_id + + service_data: dict[str, Any] = { + ATTR_ENTITY_ID: entity_id, + ATTR_MEDIA_SEARCH_QUERY: args[ATTR_MEDIA_SEARCH_QUERY], + } + if media_class := args.get("media_class"): + service_data[ATTR_MEDIA_FILTER_CLASSES] = [media_class] + + service_result = await hass.services.async_call( + DOMAIN, + SERVICE_SEARCH_MEDIA, + service_data, + context=llm_context.context, + blocking=True, + return_response=True, + ) + search_media = cast(dict[str, SearchMedia], service_result)[entity_id] + playable = [item for item in search_media.result if item.can_play] + results: list[JsonValueType] = [ + { + "title": item.title, + "media_class": item.media_class, + ATTR_MEDIA_CONTENT_TYPE: item.media_content_type, + ATTR_MEDIA_CONTENT_ID: item.media_content_id, + } + for item in playable[:MAX_SEARCH_RESULTS] + ] + if not results: + return ToolResult(data={"results": results}) + return ToolResult( + data={ + "results": results, + "instruction": ( + f"To play a result, call {MediaPlayTool.name} with its " + "media_content_id and media_content_type, and with the same " + "player_name, player_area and player_floor as this search." + ), + } + ) + + +class MediaPlayTool(Tool): + """LLM Tool that plays a media item found by the search media tool.""" + + name = "media_player__play_media" + title = "Play media" + description = ( + "Plays a media item that the search media tool returned, or a URL. " + "For a search result, pass the same player_name, player_area and " + "player_floor as the search." + ) + parameters = probatio.Schema( + { + probatio.Required( + ATTR_MEDIA_CONTENT_ID, + description="The media_content_id of the search result, or a URL", + ): cv.string, + probatio.Required( + ATTR_MEDIA_CONTENT_TYPE, + description=( + "The media_content_type of the search result. " + "Use music to play an audio URL." + ), + ): cv.string, + **TARGET_SCHEMA, + } + ) + integration = DOMAIN + + @override + async def async_call( + self, hass: HomeAssistant, tool_input: ToolInput, llm_context: LLMContext + ) -> ToolResult: + """Play a media item.""" + args = _validate_args(self.parameters, tool_input.tool_args) + entity_id = _async_match_player(hass, llm_context, args).entity_id + + await hass.services.async_call( + DOMAIN, + SERVICE_PLAY_MEDIA, + { + ATTR_ENTITY_ID: entity_id, + ATTR_MEDIA_CONTENT_ID: args[ATTR_MEDIA_CONTENT_ID], + ATTR_MEDIA_CONTENT_TYPE: args[ATTR_MEDIA_CONTENT_TYPE], + }, + context=llm_context.context, + blocking=True, + ) + return ToolResult(data={"success": True}) + @callback def async_get_tools( hass: HomeAssistant, llm_context: LLMContext, api_id: str ) -> LLMTools | None: - """Return LLM tools for the integration's intents when its domain is exposed.""" + """Return LLM tools for the integration when its domain is exposed.""" if api_id != LLM_API_ASSIST: return None if not llm_context.assistant: return None - if not any( - async_should_expose(hass, llm_context.assistant, state.entity_id) + exposed = [ + state for state in hass.states.async_all(DOMAIN) - ): + if async_should_expose(hass, llm_context.assistant, state.entity_id) + ] + if not exposed: return None tools: list[Tool] = [ @@ -86,4 +276,10 @@ def async_get_tools( for handler in intent.async_get(hass) if handler.intent_type in LLM_INTENTS ] + if any( + state.attributes.get(ATTR_SUPPORTED_FEATURES, 0) & SEARCH_PLAY_FEATURES + == SEARCH_PLAY_FEATURES + for state in exposed + ): + tools.extend([MediaSearchTool(), MediaPlayTool()]) return LLMTools(tools=tools) diff --git a/homeassistant/helpers/llm.py b/homeassistant/helpers/llm.py index b4659dc5b2ec..4c0f72aa2da6 100644 --- a/homeassistant/helpers/llm.py +++ b/homeassistant/helpers/llm.py @@ -308,6 +308,27 @@ class API(ABC): raise NotImplementedError +@callback +def async_get_match_preferences( + hass: HomeAssistant, llm_context: LLMContext +) -> intent.MatchTargetsPreferences: + """Return target match preferences for the area of the requesting device.""" + area: ar.AreaEntry | None = None + floor: fr.FloorEntry | None = None + if ( + llm_context.device_id + and (device := dr.async_get(hass).async_get(llm_context.device_id)) + and (device_area_id := dr.async_get_effective_area_id(hass, device)) + and (area := ar.async_get(hass).async_get_area(device_area_id)) + and area.floor_id + ): + floor = fr.async_get(hass).async_get_floor(area.floor_id) + return intent.MatchTargetsPreferences( + area_id=area.id if area else None, + floor_id=floor.floor_id if floor else None, + ) + + class IntentTool(Tool): """LLM Tool representing an Intent.""" @@ -357,24 +378,11 @@ class IntentTool(Tool): if not intent.is_blank_slot_value(val) } - if self.extra_slots and llm_context.device_id: - device_reg = dr.async_get(hass) - device = device_reg.async_get(llm_context.device_id) - - area: ar.AreaEntry | None = None - floor: fr.FloorEntry | None = None - if device: - area_reg = ar.async_get(hass) - if ( - device_area_id := dr.async_get_effective_area_id(hass, device) - ) and (area := area_reg.async_get_area(device_area_id)): - if area.floor_id: - floor_reg = fr.async_get(hass) - floor = floor_reg.async_get_floor(area.floor_id) - + if self.extra_slots: + preferences = async_get_match_preferences(hass, llm_context) for slot_name, slot_value in ( - ("preferred_area_id", area.id if area else None), - ("preferred_floor_id", floor.floor_id if floor else None), + ("preferred_area_id", preferences.area_id), + ("preferred_floor_id", preferences.floor_id), ): if slot_value and slot_name in self.extra_slots: slots[slot_name] = {"value": slot_value} diff --git a/tests/components/media_player/test_llm.py b/tests/components/media_player/test_llm.py index 559cd92b7f3e..5813a5c704cc 100644 --- a/tests/components/media_player/test_llm.py +++ b/tests/components/media_player/test_llm.py @@ -1,26 +1,78 @@ """Tests for the media_player LLM tools platform.""" +from typing import Any + +import probatio import pytest from homeassistant.components import llm as llm_component from homeassistant.components.homeassistant.exposed_entities import async_expose_entity -from homeassistant.components.media_player import llm as media_player_llm +from homeassistant.components.media_player import ( + DOMAIN, + SERVICE_PLAY_MEDIA, + SERVICE_SEARCH_MEDIA, + BrowseMedia, + MediaClass, + MediaPlayerEntityFeature, + MediaType, + SearchMedia, + llm as media_player_llm, +) +from homeassistant.const import ATTR_SUPPORTED_FEATURES from homeassistant.core import Context, HomeAssistant -from homeassistant.helpers import llm +from homeassistant.helpers import ( + area_registry as ar, + device_registry as dr, + entity_registry as er, + floor_registry as fr, + intent, + llm, +) from homeassistant.setup import async_setup_component +from tests.common import MockConfigEntry, async_mock_service + ENTITY_ID = "media_player.test" -TOOL_NAMES = { +SEARCH_PLAY_FEATURES = ( + MediaPlayerEntityFeature.SEARCH_MEDIA | MediaPlayerEntityFeature.PLAY_MEDIA +) +INTENT_TOOL_NAMES = { "media_player__HassMediaNext", "media_player__HassMediaPause", "media_player__HassMediaPlayerMute", "media_player__HassMediaPlayerUnmute", "media_player__HassMediaPrevious", - "media_player__HassMediaSearchAndPlay", "media_player__HassMediaUnpause", "media_player__HassSetVolume", "media_player__HassSetVolumeRelative", } +SEARCH_PLAY_TOOL_NAMES = {"media_player__search_media", "media_player__play_media"} +TOOL_NAMES = INTENT_TOOL_NAMES | SEARCH_PLAY_TOOL_NAMES + +TRACK = BrowseMedia( + title="Queen - Bohemian Rhapsody", + media_class=MediaClass.TRACK, + media_content_type=MediaType.TRACK, + media_content_id="library://track/1", + can_play=True, + can_expand=False, +) +ALBUM = BrowseMedia( + title="A Night at the Opera", + media_class=MediaClass.ALBUM, + media_content_type=MediaType.ALBUM, + media_content_id="library://album/2", + can_play=True, + can_expand=True, +) +ARTIST = BrowseMedia( + title="Queen", + media_class=MediaClass.ARTIST, + media_content_type=MediaType.ARTIST, + media_content_id="library://artist/3", + can_play=False, + can_expand=True, +) @pytest.fixture(autouse=True) @@ -30,19 +82,26 @@ async def setup_integrations(hass: HomeAssistant) -> None: assert await async_setup_component(hass, "intent", {}) assert await async_setup_component(hass, "media_player", {}) assert await async_setup_component(hass, "llm", {}) - hass.states.async_set(ENTITY_ID, "on", {"friendly_name": "Test media_player"}) + hass.states.async_set( + ENTITY_ID, + "idle", + { + "friendly_name": "Test media_player", + ATTR_SUPPORTED_FEATURES: SEARCH_PLAY_FEATURES, + }, + ) async_expose_entity(hass, "conversation", ENTITY_ID, True) await hass.async_block_till_done() -def _llm_context() -> llm.LLMContext: +def _llm_context(device_id: str | None = None) -> llm.LLMContext: """Return an LLM context for the conversation assistant.""" return llm.LLMContext( platform="test_platform", context=Context(), language="*", assistant="conversation", - device_id=None, + device_id=device_id, ) @@ -52,11 +111,25 @@ async def _tool_names(hass: HomeAssistant) -> set[str]: return {tool.name for tool in result.tools} -async def test_intent_tool_exposed(hass: HomeAssistant) -> None: - """Test the intent tool is offered for an exposed media_player entity.""" +async def _async_call_tool( + hass: HomeAssistant, + tool_name: str, + tool_args: dict[str, Any], + device_id: str | None = None, +) -> llm.ToolResult: + """Call a tool through the Assist API.""" + api = await llm.async_get_api(hass, "assist", _llm_context(device_id)) + return await api.async_call_tool( + llm.ToolInput(tool_name=tool_name, tool_args=tool_args) + ) + + +async def test_tools_exposed(hass: HomeAssistant) -> None: + """Test the tools are offered for an exposed media_player entity.""" result = await llm_component.async_get_tools(hass, _llm_context(), "assist") tools = {tool.name: tool for tool in result.tools} assert tools.keys() >= TOOL_NAMES + assert "media_player__HassMediaSearchAndPlay" not in tools control = llm.ToolAnnotations(idempotent=True, open_world=False) repeats = llm.ToolAnnotations(open_world=False) @@ -74,11 +147,6 @@ async def test_intent_tool_exposed(hass: HomeAssistant) -> None: control, ), "media_player__HassMediaPrevious": ("Previous track", "media_player", repeats), - "media_player__HassMediaSearchAndPlay": ( - "Search and play media", - "media_player", - llm.ToolAnnotations(), - ), "media_player__HassMediaUnpause": ("Resume media", "media_player", repeats), "media_player__HassSetVolume": ("Set volume", "media_player", control), "media_player__HassSetVolumeRelative": ( @@ -86,16 +154,381 @@ async def test_intent_tool_exposed(hass: HomeAssistant) -> None: "media_player", repeats, ), + "media_player__search_media": ( + "Search media", + "media_player", + llm.ToolAnnotations(read_only=True, destructive=False, idempotent=True), + ), + "media_player__play_media": ( + "Play media", + "media_player", + llm.ToolAnnotations(), + ), } -async def test_intent_tool_not_exposed(hass: HomeAssistant) -> None: - """Test the intent tool is hidden when no media_player entity is exposed.""" +async def test_tools_not_exposed(hass: HomeAssistant) -> None: + """Test the tools are hidden when no media_player entity is exposed.""" async_expose_entity(hass, "conversation", ENTITY_ID, False) assert not TOOL_NAMES & await _tool_names(hass) assert media_player_llm.async_get_tools(hass, _llm_context(), "assist") is None +@pytest.mark.parametrize( + "supported_features", + [ + pytest.param(0, id="none"), + pytest.param(MediaPlayerEntityFeature.SEARCH_MEDIA, id="search_only"), + pytest.param(MediaPlayerEntityFeature.PLAY_MEDIA, id="play_only"), + ], +) +async def test_search_play_tools_need_features( + hass: HomeAssistant, supported_features: MediaPlayerEntityFeature +) -> None: + """Test the search and play tools need a player that can search and play.""" + hass.states.async_set( + ENTITY_ID, "idle", {ATTR_SUPPORTED_FEATURES: supported_features} + ) + tool_names = await _tool_names(hass) + assert tool_names >= INTENT_TOOL_NAMES + assert not SEARCH_PLAY_TOOL_NAMES & tool_names + + async def test_no_tools_for_other_api(hass: HomeAssistant) -> None: """Test the platform returns None for an unsupported API.""" assert media_player_llm.async_get_tools(hass, _llm_context(), "other") is None + + +@pytest.mark.parametrize( + ("tool_args", "service_data"), + [ + pytest.param( + {"search_query": "queen"}, + {"entity_id": ENTITY_ID, "search_query": "queen"}, + id="query", + ), + pytest.param( + {"search_query": "queen", "media_class": "album"}, + { + "entity_id": ENTITY_ID, + "search_query": "queen", + "media_filter_classes": ["album"], + }, + id="media_class", + ), + ], +) +async def test_search_media( + hass: HomeAssistant, tool_args: dict[str, Any], service_data: dict[str, Any] +) -> None: + """Test the search tool returns the playable results of the player.""" + search_calls = async_mock_service( + hass, + DOMAIN, + SERVICE_SEARCH_MEDIA, + response={ENTITY_ID: SearchMedia(result=[TRACK, ARTIST, ALBUM])}, + ) + + result = await _async_call_tool(hass, "media_player__search_media", tool_args) + + assert result == llm.ToolResult( + data={ + "results": [ + { + "title": "Queen - Bohemian Rhapsody", + "media_class": "track", + "media_content_type": "track", + "media_content_id": "library://track/1", + }, + { + "title": "A Night at the Opera", + "media_class": "album", + "media_content_type": "album", + "media_content_id": "library://album/2", + }, + ], + "instruction": ( + "To play a result, call media_player__play_media with its " + "media_content_id and media_content_type, and with the same " + "player_name, player_area and player_floor as this search." + ), + } + ) + assert len(search_calls) == 1 + assert search_calls[0].data == service_data + + +async def test_search_media_limits_results(hass: HomeAssistant) -> None: + """Test the search tool returns at most 20 playable results.""" + tracks = [ + BrowseMedia( + title=f"Track {index}", + media_class=MediaClass.TRACK, + media_content_type=MediaType.TRACK, + media_content_id=f"library://track/{index}", + can_play=True, + can_expand=False, + ) + for index in range(25) + ] + async_mock_service( + hass, + DOMAIN, + SERVICE_SEARCH_MEDIA, + response={ENTITY_ID: SearchMedia(result=[ARTIST, *tracks])}, + ) + + result = await _async_call_tool( + hass, "media_player__search_media", {"search_query": "track"} + ) + + assert [item["title"] for item in result.data["results"]] == [ + f"Track {index}" for index in range(20) + ] + + +async def test_search_media_no_results(hass: HomeAssistant) -> None: + """Test the search tool returns an empty list when nothing matches.""" + async_mock_service( + hass, + DOMAIN, + SERVICE_SEARCH_MEDIA, + response={ENTITY_ID: SearchMedia(result=[])}, + ) + + result = await _async_call_tool( + hass, "media_player__search_media", {"search_query": "nothing"} + ) + + assert result == llm.ToolResult(data={"results": []}) + + +async def test_search_media_invalid_media_class(hass: HomeAssistant) -> None: + """Test the search tool rejects an unknown media class.""" + search_calls = async_mock_service(hass, DOMAIN, SERVICE_SEARCH_MEDIA) + + with pytest.raises(probatio.Invalid): + await _async_call_tool( + hass, + "media_player__search_media", + {"search_query": "queen", "media_class": "invalid"}, + ) + assert not search_calls + + +@pytest.mark.parametrize( + "media", + [ + pytest.param( + {"media_content_id": "library://album/2", "media_content_type": "album"}, + id="search_result", + ), + pytest.param( + { + "media_content_id": "https://example.com/stream.mp3", + "media_content_type": "music", + }, + id="url", + ), + ], +) +async def test_play_media(hass: HomeAssistant, media: dict[str, str]) -> None: + """Test the play tool plays the chosen item on the player.""" + play_calls = async_mock_service(hass, DOMAIN, SERVICE_PLAY_MEDIA) + + result = await _async_call_tool( + hass, "media_player__play_media", {**media, "player_name": "Test media_player"} + ) + + assert result == llm.ToolResult(data={"success": True}) + assert len(play_calls) == 1 + assert play_calls[0].data == {"entity_id": ENTITY_ID, **media} + + +async def test_blank_target_values_omitted(hass: HomeAssistant) -> None: + """Test the tools treat blank target values as omitted.""" + search_calls = async_mock_service( + hass, + DOMAIN, + SERVICE_SEARCH_MEDIA, + response={ENTITY_ID: SearchMedia(result=[])}, + ) + play_calls = async_mock_service(hass, DOMAIN, SERVICE_PLAY_MEDIA) + blank_target = {"player_name": "", "player_area": " ", "player_floor": None} + + await _async_call_tool( + hass, + "media_player__search_media", + {"search_query": "queen", "media_class": "", **blank_target}, + ) + await _async_call_tool( + hass, + "media_player__play_media", + { + "media_content_id": "library://album/2", + "media_content_type": "album", + **blank_target, + }, + ) + + assert search_calls[0].data == {"entity_id": ENTITY_ID, "search_query": "queen"} + assert play_calls[0].data["entity_id"] == ENTITY_ID + + +@pytest.mark.parametrize( + ("tool_name", "tool_args"), + [ + pytest.param( + "media_player__search_media", {"search_query": "queen"}, id="search" + ), + pytest.param( + "media_player__play_media", + {"media_content_id": "library://album/2", "media_content_type": "album"}, + id="play", + ), + ], +) +async def test_player_without_features_not_matched( + hass: HomeAssistant, tool_name: str, tool_args: dict[str, Any] +) -> None: + """Test the tools only target a player that can search and play.""" + hass.states.async_set( + "media_player.play_only", + "idle", + { + "friendly_name": "Play only", + ATTR_SUPPORTED_FEATURES: MediaPlayerEntityFeature.PLAY_MEDIA, + }, + ) + async_expose_entity(hass, "conversation", "media_player.play_only", True) + search_calls = async_mock_service(hass, DOMAIN, SERVICE_SEARCH_MEDIA) + play_calls = async_mock_service(hass, DOMAIN, SERVICE_PLAY_MEDIA) + + with pytest.raises(intent.MatchFailedError): + await _async_call_tool( + hass, tool_name, {**tool_args, "player_name": "Play only"} + ) + assert not search_calls + assert not play_calls + + +@pytest.mark.parametrize( + "target", + [ + pytest.param({"player_name": "Kitchen speaker"}, id="player_name"), + pytest.param({"player_area": "Kitchen"}, id="player_area"), + pytest.param({"player_floor": "Ground floor"}, id="player_floor"), + ], +) +async def test_play_media_target( + hass: HomeAssistant, + area_registry: ar.AreaRegistry, + entity_registry: er.EntityRegistry, + floor_registry: fr.FloorRegistry, + target: dict[str, str], +) -> None: + """Test the player target arguments select the media player.""" + ground_floor = floor_registry.async_create("Ground floor") + upstairs = floor_registry.async_create("Upstairs") + kitchen = area_registry.async_create("Kitchen", floor_id=ground_floor.floor_id) + bedroom = area_registry.async_create("Bedroom", floor_id=upstairs.floor_id) + async_expose_entity(hass, "conversation", ENTITY_ID, False) + + for area in (kitchen, bedroom): + entry = entity_registry.async_get_or_create( + DOMAIN, + "test", + area.id, + suggested_object_id=area.id, + original_name=f"{area.name} speaker", + ) + entity_registry.async_update_entity(entry.entity_id, area_id=area.id) + hass.states.async_set( + entry.entity_id, + "idle", + { + "friendly_name": f"{area.name} speaker", + ATTR_SUPPORTED_FEATURES: SEARCH_PLAY_FEATURES, + }, + ) + async_expose_entity(hass, "conversation", entry.entity_id, True) + + play_calls = async_mock_service(hass, DOMAIN, SERVICE_PLAY_MEDIA) + + await _async_call_tool( + hass, + "media_player__play_media", + { + "media_content_id": "library://track/1", + "media_content_type": "track", + **target, + }, + ) + + assert play_calls[0].data["entity_id"] == "media_player.kitchen" + + +async def test_search_and_play_use_device_area( + hass: HomeAssistant, + area_registry: ar.AreaRegistry, + device_registry: dr.DeviceRegistry, + entity_registry: er.EntityRegistry, +) -> None: + """Test search and play target the player in the area of the device.""" + kitchen = area_registry.async_create("Kitchen") + bedroom = area_registry.async_create("Bedroom") + + for area in (kitchen, bedroom): + entry = entity_registry.async_get_or_create( + DOMAIN, "test", area.id, suggested_object_id=area.id + ) + entity_registry.async_update_entity(entry.entity_id, area_id=area.id) + hass.states.async_set( + entry.entity_id, + "idle", + { + "friendly_name": f"{area.name} speaker", + ATTR_SUPPORTED_FEATURES: SEARCH_PLAY_FEATURES, + }, + ) + async_expose_entity(hass, "conversation", entry.entity_id, True) + + config_entry = MockConfigEntry() + config_entry.add_to_hass(hass) + satellite = device_registry.async_get_or_create( + config_entry_id=config_entry.entry_id, + connections={(dr.CONNECTION_NETWORK_MAC, "12:34:56:78:90:ab")}, + ) + device_registry.async_update_device(satellite.id, area_id=bedroom.id) + + search_calls = async_mock_service( + hass, + DOMAIN, + SERVICE_SEARCH_MEDIA, + response={"media_player.bedroom": SearchMedia(result=[TRACK])}, + ) + play_calls = async_mock_service(hass, DOMAIN, SERVICE_PLAY_MEDIA) + + result = await _async_call_tool( + hass, + "media_player__search_media", + {"search_query": "bohemian rhapsody"}, + device_id=satellite.id, + ) + item = result.data["results"][0] + await _async_call_tool( + hass, + "media_player__play_media", + { + "media_content_id": item["media_content_id"], + "media_content_type": item["media_content_type"], + }, + device_id=satellite.id, + ) + + assert search_calls[0].data["entity_id"] == "media_player.bedroom" + assert play_calls[0].data == { + "entity_id": "media_player.bedroom", + "media_content_id": "library://track/1", + "media_content_type": "track", + }