mirror of
https://github.com/home-assistant/core.git
synced 2026-10-06 22:38:02 -04:00
Require an integration on an LLM tool (#182714)
Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude
parent
7398141b4c
commit
fac1ed7d4c
@@ -15,6 +15,7 @@ from homeassistant.helpers.llm import (
|
||||
LLMContext,
|
||||
Tool,
|
||||
async_register_api,
|
||||
report_untagged_tool,
|
||||
selector_serializer,
|
||||
)
|
||||
from homeassistant.helpers.typing import ConfigType
|
||||
@@ -91,7 +92,7 @@ async def async_get_tools(
|
||||
continue
|
||||
if result is None:
|
||||
continue
|
||||
_async_report_unprefixed_tools(hass, domain, result.tools)
|
||||
_async_report_tool_issues(hass, domain, result.tools)
|
||||
tools.extend(result.tools)
|
||||
if result.prompt:
|
||||
prompts.append(result.prompt)
|
||||
@@ -99,12 +100,18 @@ async def async_get_tools(
|
||||
|
||||
|
||||
@callback
|
||||
def _async_report_unprefixed_tools(
|
||||
def _async_report_tool_issues(
|
||||
hass: HomeAssistant, domain: str, tools: list[Tool]
|
||||
) -> None:
|
||||
"""Report tools that are not prefixed with the domain offering them."""
|
||||
"""Report tools that do not follow the current requirements."""
|
||||
prefix = f"{domain}__"
|
||||
unprefixed = [tool.name for tool in tools if not tool.name.startswith(prefix)]
|
||||
unprefixed: list[str] = []
|
||||
for tool in tools:
|
||||
if not tool.name.startswith(prefix):
|
||||
unprefixed.append(tool.name)
|
||||
if tool.integration is None:
|
||||
report_untagged_tool(tool, domain)
|
||||
|
||||
if not unprefixed:
|
||||
return
|
||||
|
||||
|
||||
@@ -42,6 +42,8 @@ APIS_CACHE: HassKey[dict[str, API]] = HassKey("llm_apis")
|
||||
|
||||
LLM_API_ASSIST = "assist"
|
||||
|
||||
TOOL_INTEGRATION_BREAKS_IN_HA_VERSION = "2027.10"
|
||||
|
||||
DATE_TIME_PROMPT = (
|
||||
'Current time is {{ now().strftime("%H:%M:%S") }}. '
|
||||
'Today\'s date is {{ now().strftime("%Y-%m-%d") }}.\n'
|
||||
@@ -210,6 +212,19 @@ class APIInstance:
|
||||
tools: list[Tool]
|
||||
custom_serializer: Callable[[Any], Any] | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Report a tool that does not record the integration providing it."""
|
||||
for tool in self.tools:
|
||||
if tool.integration is not None:
|
||||
continue
|
||||
# A tool class outside an integration, such as a shared helper tool,
|
||||
# belongs to whichever integration provides the API.
|
||||
domain = _tool_integration_domain(tool) or _integration_domain(
|
||||
type(self.api).__module__
|
||||
)
|
||||
if domain is not None:
|
||||
report_untagged_tool(tool, domain)
|
||||
|
||||
async def async_call_tool(self, tool_input: ToolInput) -> ToolResult:
|
||||
"""Call a LLM tool, validate args and return the response."""
|
||||
from homeassistant.components.conversation import ( # noqa: PLC0415
|
||||
@@ -244,11 +259,35 @@ class APIInstance:
|
||||
return ToolResult(data=result)
|
||||
|
||||
|
||||
@callback
|
||||
def report_untagged_tool(tool: Tool, domain: str) -> None:
|
||||
"""Report a tool that does not record the integration providing it."""
|
||||
wrapped = [tool]
|
||||
while isinstance(wrapped[-1], NamespacedTool):
|
||||
wrapped.append(wrapped[-1].tool)
|
||||
frame.report_usage(
|
||||
f"provides the LLM tool {wrapped[-1].name} without an integration",
|
||||
breaks_in_ha_version=TOOL_INTEGRATION_BREAKS_IN_HA_VERSION,
|
||||
core_behavior=frame.ReportBehavior.ERROR,
|
||||
core_integration_behavior=frame.ReportBehavior.ERROR,
|
||||
custom_integration_behavior=frame.ReportBehavior.LOG,
|
||||
integration_domain=domain,
|
||||
)
|
||||
# Record the domain on the tool and every wrapper around it, so it carries
|
||||
# the integration until the requirement is enforced.
|
||||
for entry in wrapped:
|
||||
entry.integration = domain
|
||||
|
||||
|
||||
def _tool_integration_domain(tool: Tool) -> str | None:
|
||||
"""Return the domain of the integration that provides the tool."""
|
||||
while isinstance(tool, NamespacedTool):
|
||||
tool = tool.tool
|
||||
module = type(tool).__module__
|
||||
return _integration_domain(type(tool).__module__)
|
||||
|
||||
|
||||
def _integration_domain(module: str) -> str | None:
|
||||
"""Return the domain of the integration that defines the module."""
|
||||
for prefix in ("custom_components.", "homeassistant.components."):
|
||||
if module.startswith(prefix):
|
||||
return module.removeprefix(prefix).partition(".")[0]
|
||||
|
||||
@@ -141,6 +141,7 @@ async def test_multiple_llm_apis(
|
||||
|
||||
name = "test_tool"
|
||||
description = "Test function"
|
||||
integration = "test"
|
||||
parameters = probatio.Schema(
|
||||
{probatio.Optional("param1", description="Test parameters"): str}
|
||||
)
|
||||
|
||||
@@ -17,10 +17,11 @@ from tests.common import mock_platform
|
||||
class _StubTool(llm.Tool):
|
||||
"""Minimal tool for registry tests."""
|
||||
|
||||
def __init__(self, name: str) -> None:
|
||||
def __init__(self, name: str, integration: str | None = "test") -> None:
|
||||
"""Initialize the stub tool."""
|
||||
self.name = name
|
||||
self.description = f"{name} description"
|
||||
self.integration = integration
|
||||
|
||||
async def async_call(
|
||||
self,
|
||||
@@ -188,6 +189,71 @@ async def test_get_tools_reports_unprefixed_tool_names(
|
||||
assert record.levelno == expected_level
|
||||
|
||||
|
||||
async def test_get_tools_untagged_tool_raises_for_core(
|
||||
hass: HomeAssistant,
|
||||
llm_context: llm.LLMContext,
|
||||
) -> None:
|
||||
"""Test a core integration must record the integration on its tools."""
|
||||
tools = [_StubTool("test__untagged", integration=None)]
|
||||
_mock_tools_platform(hass, "test", LLMTools(tools=tools))
|
||||
|
||||
assert await async_setup_component(hass, "llm", {})
|
||||
|
||||
with (
|
||||
patch.object(frame, "_REPORTED_INTEGRATIONS", set()),
|
||||
pytest.raises(
|
||||
RuntimeError,
|
||||
match="provides the LLM tool test__untagged without an integration",
|
||||
),
|
||||
):
|
||||
await async_get_tools(hass, llm_context, "assist")
|
||||
|
||||
|
||||
async def test_get_tools_untagged_tool_reported_for_custom(
|
||||
hass: HomeAssistant,
|
||||
llm_context: llm.LLMContext,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
"""Test a custom integration is warned about tools without an integration."""
|
||||
tools = [_StubTool("test__untagged", integration=None)]
|
||||
_mock_tools_platform(hass, "test", LLMTools(tools=tools), built_in=False)
|
||||
|
||||
assert await async_setup_component(hass, "llm", {})
|
||||
|
||||
with patch.object(frame, "_REPORTED_INTEGRATIONS", set()):
|
||||
result = await async_get_tools(hass, llm_context, "assist")
|
||||
|
||||
# The tool is still returned until the requirement starts to fail.
|
||||
assert "test__untagged" in [tool.name for tool in result.tools]
|
||||
assert (
|
||||
"Detected that custom integration 'test' provides the LLM tool test__untagged "
|
||||
"without an integration. This will stop working in Home Assistant 2027.10"
|
||||
in caplog.text
|
||||
)
|
||||
# The platform domain is recorded on the tool, so it is not reported again.
|
||||
assert tools[0].integration == "test"
|
||||
|
||||
|
||||
async def test_get_tools_untagged_wrapped_tool_reported_once(
|
||||
hass: HomeAssistant,
|
||||
llm_context: llm.LLMContext,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
"""Test a tool wrapper is tagged along with the tool it wraps."""
|
||||
wrapper = llm.NamespacedTool("test", _StubTool("untagged", integration=None))
|
||||
_mock_tools_platform(hass, "test", LLMTools(tools=[wrapper]), built_in=False)
|
||||
|
||||
assert await async_setup_component(hass, "llm", {})
|
||||
|
||||
with patch.object(frame, "_REPORTED_INTEGRATIONS", set()):
|
||||
await async_get_tools(hass, llm_context, "assist")
|
||||
|
||||
# The wrapper is tagged too, so the API instance does not report it again.
|
||||
assert wrapper.integration == "test"
|
||||
assert wrapper.tool.integration == "test"
|
||||
assert caplog.text.count("without an integration") == 1
|
||||
|
||||
|
||||
async def test_get_tools_prefixed_tool_names_not_reported(
|
||||
hass: HomeAssistant,
|
||||
llm_context: llm.LLMContext,
|
||||
|
||||
@@ -86,6 +86,7 @@ class _StubTool(llm.Tool):
|
||||
"""Minimal tool with a configurable parameter schema."""
|
||||
|
||||
name = "test_tool"
|
||||
integration = "test"
|
||||
|
||||
def __init__(self, parameters: probatio.Schema) -> None:
|
||||
"""Initialize the stub tool."""
|
||||
|
||||
@@ -230,6 +230,100 @@ async def test_call_tool_deprecated_json_object_custom_integration(
|
||||
assert "returns a JSON object from a tool" in caplog.text
|
||||
|
||||
|
||||
def _untagged_tool(module: str) -> llm.Tool:
|
||||
"""Return a tool that does not record the integration providing it."""
|
||||
|
||||
class UntaggedTool(llm.Tool):
|
||||
"""Tool that declares no integration."""
|
||||
|
||||
name = "test_tool"
|
||||
|
||||
async def async_call(
|
||||
self, hass: HomeAssistant, tool_input: llm.ToolInput, _: llm.LLMContext
|
||||
) -> llm.ToolResult:
|
||||
return llm.ToolResult(data={})
|
||||
|
||||
# The tool is reported against the integration its class comes from.
|
||||
UntaggedTool.__module__ = module
|
||||
return UntaggedTool()
|
||||
|
||||
|
||||
async def test_api_instance_reports_untagged_tool_for_custom_integration(
|
||||
hass: HomeAssistant,
|
||||
llm_context: llm.LLMContext,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
"""Test a custom integration is warned about a tool without an integration."""
|
||||
mock_integration(hass, MockModule("my_custom"), built_in=False)
|
||||
tool = _untagged_tool("custom_components.my_custom.llm")
|
||||
|
||||
llm.APIInstance(MyAPI(hass=hass, id="test", name="Test"), "", llm_context, [tool])
|
||||
|
||||
assert "provides the LLM tool test_tool without an integration" in caplog.text
|
||||
|
||||
|
||||
async def test_api_instance_raises_untagged_tool_for_core_integration(
|
||||
hass: HomeAssistant,
|
||||
llm_context: llm.LLMContext,
|
||||
) -> None:
|
||||
"""Test a core integration must record the integration on its tools."""
|
||||
mock_integration(hass, MockModule("my_core"))
|
||||
tool = _untagged_tool("homeassistant.components.my_core.llm")
|
||||
|
||||
with pytest.raises(
|
||||
RuntimeError, match="provides the LLM tool test_tool without an integration"
|
||||
):
|
||||
llm.APIInstance(
|
||||
MyAPI(hass=hass, id="test", name="Test"), "", llm_context, [tool]
|
||||
)
|
||||
|
||||
|
||||
async def test_api_instance_reports_untagged_tool_from_the_api(
|
||||
hass: HomeAssistant,
|
||||
llm_context: llm.LLMContext,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
"""Test a tool defined outside an integration is reported against its API."""
|
||||
mock_integration(hass, MockModule("my_custom"), built_in=False)
|
||||
tool = _untagged_tool("homeassistant.helpers.llm")
|
||||
|
||||
class CustomAPI(MyAPI):
|
||||
"""API provided by a custom integration."""
|
||||
|
||||
CustomAPI.__module__ = "custom_components.my_custom.llm_api"
|
||||
llm.APIInstance(
|
||||
CustomAPI(hass=hass, id="test", name="Test"), "", llm_context, [tool]
|
||||
)
|
||||
|
||||
assert (
|
||||
"custom integration 'my_custom' provides the LLM tool test_tool without an "
|
||||
"integration" in caplog.text
|
||||
)
|
||||
# The tool carries the domain until the requirement is enforced.
|
||||
assert tool.integration == "my_custom"
|
||||
|
||||
|
||||
async def test_merged_api_reports_untagged_tool_once(
|
||||
hass: HomeAssistant,
|
||||
llm_context: llm.LLMContext,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
"""Test a merged API reports the wrapped tool under its own name."""
|
||||
mock_integration(hass, MockModule("my_custom"), built_in=False)
|
||||
|
||||
api = MyAPI(hass=hass, id="api-1", name="API 1")
|
||||
api.tools = [_untagged_tool("custom_components.my_custom.llm")]
|
||||
llm.async_register_api(hass, api)
|
||||
other = MyAPI(hass=hass, id="api-2", name="API 2")
|
||||
llm.async_register_api(hass, other)
|
||||
|
||||
await llm.async_get_api(hass, ["api-1", "api-2"], llm_context)
|
||||
|
||||
assert "provides the LLM tool test_tool without an integration" in caplog.text
|
||||
# The wrapper reports the tool it wraps, so the report is not repeated.
|
||||
assert caplog.text.count("without an integration") == 1
|
||||
|
||||
|
||||
def test_tool_metadata_defaults() -> None:
|
||||
"""Test a tool that declares no metadata is taken to be unsafe."""
|
||||
|
||||
@@ -313,7 +407,7 @@ async def test_intent_tool_omits_blank_arguments(
|
||||
probatio.Optional("enabled"): cv.boolean,
|
||||
}
|
||||
|
||||
intent_tool = llm.IntentTool("test_intent", MyIntentHandler())
|
||||
intent_tool = llm.IntentTool("test_intent", MyIntentHandler(), integration="test")
|
||||
tool: llm.Tool = (
|
||||
llm.NamespacedTool("test_api", intent_tool) if namespaced else intent_tool
|
||||
)
|
||||
@@ -379,7 +473,7 @@ async def test_assist_api(
|
||||
|
||||
intent_handler = MyIntentHandler()
|
||||
|
||||
tool = llm.IntentTool("test_intent", intent_handler)
|
||||
tool = llm.IntentTool("test_intent", intent_handler, integration="test")
|
||||
assert tool.name == "test_intent"
|
||||
assert tool.description == "Execute Home Assistant test_intent intent"
|
||||
assert tool.parameters == probatio.Schema(
|
||||
@@ -541,7 +635,7 @@ async def test_assist_api_description(
|
||||
intent_type = "test_intent"
|
||||
description = "my intent handler"
|
||||
|
||||
tool = llm.IntentTool("test_intent", MyIntentHandler())
|
||||
tool = llm.IntentTool("test_intent", MyIntentHandler(), integration="test")
|
||||
assert tool.name == "test_intent"
|
||||
assert tool.description == "my intent handler"
|
||||
|
||||
@@ -1483,6 +1577,7 @@ async def test_merged_api(hass: HomeAssistant, llm_context: llm.LLMContext) -> N
|
||||
def __init__(self, name: str, description: str) -> None:
|
||||
self.name = name
|
||||
self.description = description
|
||||
self.integration = "test"
|
||||
|
||||
async def async_call(
|
||||
self, hass: HomeAssistant, tool_input: llm.ToolInput, _: llm.LLMContext
|
||||
|
||||
Reference in New Issue
Block a user