diff --git a/homeassistant/components/media_player/llm.py b/homeassistant/components/media_player/llm.py index fda4a790624b..6e8e29ebf155 100644 --- a/homeassistant/components/media_player/llm.py +++ b/homeassistant/components/media_player/llm.py @@ -131,13 +131,31 @@ class MediaSearchTool(Tool): name = "media_player__search_media" title = "Search media" - description = "Searches a media player for media and returns the playable items." + description = ( + "Searches a media player for media. " + "Can also search inside a result of an earlier search, " + "such as an artist or an album." + ) 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( + "within_media_content_id", + description=( + "The media_content_id of a result to search inside, " + "such as an artist. Copy it from that result." + ), + ): cv.string, + probatio.Optional( + "within_media_content_type", + description=( + "The media_content_type of the result to search inside. " + "Copy it from that result." + ), + ): cv.string, **TARGET_SCHEMA, } ) @@ -156,6 +174,10 @@ class MediaSearchTool(Tool): ATTR_ENTITY_ID: entity_id, ATTR_MEDIA_SEARCH_QUERY: args[ATTR_MEDIA_SEARCH_QUERY], } + if within_id := args.get("within_media_content_id"): + service_data[ATTR_MEDIA_CONTENT_ID] = within_id + if within_type := args.get("within_media_content_type"): + service_data[ATTR_MEDIA_CONTENT_TYPE] = within_type service_result = await hass.services.async_call( DOMAIN, @@ -166,15 +188,16 @@ class MediaSearchTool(Tool): 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, + "can_play": item.can_play, + "can_search": item.can_search, } - for item in playable[:MAX_SEARCH_RESULTS] + for item in search_media.result[:MAX_SEARCH_RESULTS] ] if not results: return ToolResult(data={"results": results}) @@ -182,9 +205,13 @@ class MediaSearchTool(Tool): 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." + f"To play a result that can_play, call {MediaPlayTool.name} " + "with its media_content_id and media_content_type. " + f"To search inside a result that can_search, call {self.name} " + "again with its media_content_id and media_content_type as " + "within_media_content_id and within_media_content_type. " + "Pass the same player_name, player_area and player_floor " + "as this search." ), } ) diff --git a/tests/components/media_player/test_llm.py b/tests/components/media_player/test_llm.py index f0a052e3135b..e91aa0e2a873 100644 --- a/tests/components/media_player/test_llm.py +++ b/tests/components/media_player/test_llm.py @@ -63,6 +63,7 @@ ALBUM = BrowseMedia( media_content_id="library://album/2", can_play=True, can_expand=True, + can_search=True, ) ARTIST = BrowseMedia( title="Queen", @@ -71,6 +72,7 @@ ARTIST = BrowseMedia( media_content_id="library://artist/3", can_play=False, can_expand=True, + can_search=True, ) @@ -199,7 +201,7 @@ async def test_no_tools_for_other_api(hass: HomeAssistant) -> None: async def test_search_media(hass: HomeAssistant) -> None: - """Test the search tool returns the playable results of the player.""" + """Test the search tool returns the results of the player.""" search_calls = async_mock_service( hass, DOMAIN, @@ -219,18 +221,34 @@ async def test_search_media(hass: HomeAssistant) -> None: "media_class": "track", "media_content_type": "track", "media_content_id": "library://track/1", + "can_play": True, + "can_search": False, + }, + { + "title": "Queen", + "media_class": "artist", + "media_content_type": "artist", + "media_content_id": "library://artist/3", + "can_play": False, + "can_search": True, }, { "title": "A Night at the Opera", "media_class": "album", "media_content_type": "album", "media_content_id": "library://album/2", + "can_play": True, + "can_search": True, }, ], "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." + "To play a result that can_play, call media_player__play_media " + "with its media_content_id and media_content_type. " + "To search inside a result that can_search, call " + "media_player__search_media again with its media_content_id and " + "media_content_type as within_media_content_id and " + "within_media_content_type. Pass the same player_name, " + "player_area and player_floor as this search." ), } ) @@ -239,7 +257,7 @@ async def test_search_media(hass: HomeAssistant) -> None: async def test_search_media_limits_results(hass: HomeAssistant) -> None: - """Test the search tool returns at most 35 playable results.""" + """Test the search tool returns at most 35 results.""" tracks = [ BrowseMedia( title=f"Track {index}", @@ -255,7 +273,7 @@ async def test_search_media_limits_results(hass: HomeAssistant) -> None: hass, DOMAIN, SERVICE_SEARCH_MEDIA, - response={ENTITY_ID: SearchMedia(result=[ARTIST, *tracks])}, + response={ENTITY_ID: SearchMedia(result=tracks)}, ) result = await _async_call_tool( @@ -267,6 +285,48 @@ async def test_search_media_limits_results(hass: HomeAssistant) -> None: ] +@pytest.mark.parametrize( + ("within", "service_data"), + [ + pytest.param( + { + "within_media_content_id": "library://artist/3", + "within_media_content_type": "artist", + }, + {"media_content_id": "library://artist/3", "media_content_type": "artist"}, + id="id_and_type", + ), + pytest.param( + {"within_media_content_id": "library://artist/3"}, + {"media_content_id": "library://artist/3"}, + id="id_only", + ), + ], +) +async def test_search_media_within_result( + hass: HomeAssistant, within: dict[str, str], service_data: dict[str, str] +) -> None: + """Test the search tool searches inside a result of an earlier search.""" + search_calls = async_mock_service( + hass, + DOMAIN, + SERVICE_SEARCH_MEDIA, + response={ENTITY_ID: SearchMedia(result=[TRACK])}, + ) + + await _async_call_tool( + hass, + "media_player__search_media", + {"search_query": "bohemian rhapsody", **within}, + ) + + assert search_calls[0].data == { + "entity_id": ENTITY_ID, + "search_query": "bohemian rhapsody", + **service_data, + } + + async def test_search_media_no_results(hass: HomeAssistant) -> None: """Test the search tool returns an empty list when nothing matches.""" async_mock_service( @@ -326,7 +386,12 @@ async def test_blank_target_values_omitted(hass: HomeAssistant) -> None: await _async_call_tool( hass, "media_player__search_media", - {"search_query": "queen", **blank_target}, + { + "search_query": "queen", + "within_media_content_id": "", + "within_media_content_type": " ", + **blank_target, + }, ) await _async_call_tool( hass,