mirror of
https://github.com/home-assistant/core.git
synced 2026-10-06 14:29:21 -04:00
Add dedicated media search and play LLM tools (#183192)
Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude
parent
8b96066fd9
commit
d6383e3629
@@ -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)
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user