mirror of
https://github.com/home-assistant/core.git
synced 2026-10-06 22:38:02 -04:00
Handle Anthropic stop reasons (#183384)
This commit is contained in:
@@ -45,6 +45,7 @@ from anthropic.types import (
|
||||
ServerToolUseBlock,
|
||||
ServerToolUseBlockParam,
|
||||
SignatureDelta,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
TextBlockParam,
|
||||
TextCitation,
|
||||
@@ -529,6 +530,8 @@ class AnthropicDeltaStream:
|
||||
"""Initialize the delta stream."""
|
||||
self._chat_log: conversation.ChatLog = chat_log
|
||||
self._stream: AsyncStream[MessageStreamEvent] = stream
|
||||
self.stop_reason: StopReason | None = None
|
||||
self.tool_args_error: json.JSONDecodeError | None = None
|
||||
|
||||
self._buffer: deque[
|
||||
conversation.AssistantContentDeltaDict
|
||||
@@ -823,34 +826,42 @@ class AnthropicDeltaStream:
|
||||
|
||||
def on_content_block_stop_event(self, index: int) -> None:
|
||||
"""Handle RawContentBlockStopEvent."""
|
||||
if self._current_tool_block is not None:
|
||||
tool_args = (
|
||||
json.loads(self._current_tool_args) if self._current_tool_args else {}
|
||||
)
|
||||
self._current_tool_block["input"] |= tool_args
|
||||
self._buffer.append(
|
||||
{
|
||||
"tool_calls": [
|
||||
llm.ToolInput(
|
||||
id=self._current_tool_block["id"],
|
||||
tool_name=self._current_tool_block["name"],
|
||||
tool_args=self._current_tool_block["input"],
|
||||
external=self._current_tool_block["type"]
|
||||
== "server_tool_use",
|
||||
)
|
||||
]
|
||||
}
|
||||
)
|
||||
self._current_tool_block = None
|
||||
if (tool_block := self._current_tool_block) is None:
|
||||
return
|
||||
self._current_tool_block = None
|
||||
tool_args_json = self._current_tool_args
|
||||
self._current_tool_args = ""
|
||||
if self.tool_args_error is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
tool_args = json.loads(tool_args_json) if tool_args_json else {}
|
||||
except json.JSONDecodeError as err:
|
||||
# Wait for the stop reason to distinguish truncation from invalid input.
|
||||
self.tool_args_error = err
|
||||
return
|
||||
|
||||
tool_block["input"] |= tool_args
|
||||
self._buffer.append(
|
||||
{
|
||||
"tool_calls": [
|
||||
llm.ToolInput(
|
||||
id=tool_block["id"],
|
||||
tool_name=tool_block["name"],
|
||||
tool_args=tool_block["input"],
|
||||
external=tool_block["type"] == "server_tool_use",
|
||||
)
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
def on_message_delta_event(self, delta: Delta, usage: MessageDeltaUsage) -> None:
|
||||
"""Handle RawMessageDeltaEvent."""
|
||||
self._chat_log.async_trace(self._create_token_stats(self._input_usage, usage))
|
||||
self._content_details.container = delta.container
|
||||
if delta.stop_reason == "refusal":
|
||||
raise HomeAssistantError(
|
||||
translation_domain=DOMAIN, translation_key="api_refusal"
|
||||
)
|
||||
if delta.container is not None:
|
||||
self._content_details.container = delta.container
|
||||
if delta.stop_reason is not None:
|
||||
self.stop_reason = delta.stop_reason
|
||||
|
||||
def on_message_stop_event(self) -> None:
|
||||
"""Handle RawMessageStopEvent."""
|
||||
@@ -1129,16 +1140,17 @@ class AnthropicBaseLLMEntity(CoordinatorEntity[AnthropicCoordinator]):
|
||||
client = coordinator.client
|
||||
|
||||
# To prevent infinite loops, we limit the number of iterations
|
||||
for _iteration in range(max_iterations):
|
||||
for iteration in range(max_iterations):
|
||||
try:
|
||||
stream = await client.messages.create(**model_args)
|
||||
delta_stream = AnthropicDeltaStream(chat_log, stream)
|
||||
|
||||
new_messages, model_args["container"] = _convert_content(
|
||||
[
|
||||
content
|
||||
async for content in chat_log.async_add_delta_content_stream(
|
||||
self.entity_id,
|
||||
AnthropicDeltaStream(chat_log, stream),
|
||||
delta_stream,
|
||||
)
|
||||
]
|
||||
)
|
||||
@@ -1174,6 +1186,34 @@ class AnthropicBaseLLMEntity(CoordinatorEntity[AnthropicCoordinator]):
|
||||
},
|
||||
) from err
|
||||
|
||||
if (stop_reason := delta_stream.stop_reason) in (
|
||||
"refusal",
|
||||
"max_tokens",
|
||||
"model_context_window_exceeded",
|
||||
"stop_sequence",
|
||||
) or (stop_reason == "pause_turn" and iteration == max_iterations - 1):
|
||||
coordinator.async_set_updated_data(coordinator.data)
|
||||
raise HomeAssistantError(
|
||||
translation_domain=DOMAIN,
|
||||
translation_key={
|
||||
"refusal": "api_refusal",
|
||||
"max_tokens": "response_max_tokens",
|
||||
"model_context_window_exceeded": "response_context_window_exceeded",
|
||||
"stop_sequence": "response_stop_sequence",
|
||||
"pause_turn": "response_incomplete",
|
||||
}[stop_reason],
|
||||
)
|
||||
|
||||
if delta_stream.tool_args_error is not None:
|
||||
coordinator.async_set_updated_data(coordinator.data)
|
||||
raise HomeAssistantError(
|
||||
translation_domain=DOMAIN,
|
||||
translation_key="tool_args_parse_error",
|
||||
) from delta_stream.tool_args_error
|
||||
|
||||
if stop_reason == "pause_turn":
|
||||
continue
|
||||
|
||||
if not chat_log.unresponded_tool_results:
|
||||
coordinator.async_set_updated_data(coordinator.data)
|
||||
break
|
||||
|
||||
@@ -181,15 +181,30 @@
|
||||
"json_parse_error": {
|
||||
"message": "Error with Claude structured response."
|
||||
},
|
||||
"response_context_window_exceeded": {
|
||||
"message": "Claude reached the context window limit before completing the response."
|
||||
},
|
||||
"response_incomplete": {
|
||||
"message": "Claude could not complete the response within the allowed number of requests."
|
||||
},
|
||||
"response_max_tokens": {
|
||||
"message": "Claude reached the output token limit before completing the response."
|
||||
},
|
||||
"response_not_found": {
|
||||
"message": "Last content in chat log is not an AssistantContent."
|
||||
},
|
||||
"response_stop_sequence": {
|
||||
"message": "Claude stopped after encountering a stop sequence."
|
||||
},
|
||||
"subentry_not_found": {
|
||||
"message": "Subentry not found."
|
||||
},
|
||||
"system_message_not_found": {
|
||||
"message": "First message must be a system message."
|
||||
},
|
||||
"tool_args_parse_error": {
|
||||
"message": "Claude returned invalid tool arguments."
|
||||
},
|
||||
"unexpected_chat_log_content": {
|
||||
"message": "Unexpected content type in chat log: {type}."
|
||||
},
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from collections.abc import AsyncGenerator, Generator, Iterable
|
||||
import datetime
|
||||
from typing import Unpack
|
||||
from unittest.mock import DEFAULT, AsyncMock, patch
|
||||
|
||||
from anthropic.pagination import AsyncPage
|
||||
@@ -18,6 +19,7 @@ from anthropic.types import (
|
||||
ToolUseBlock,
|
||||
Usage,
|
||||
)
|
||||
from anthropic.types.message_create_params import MessageCreateParamsStreaming
|
||||
from anthropic.types.raw_message_delta_event import Delta
|
||||
import pytest
|
||||
|
||||
@@ -120,16 +122,20 @@ def mock_setup_entry() -> Generator[AsyncMock]:
|
||||
def mock_create_stream() -> Generator[AsyncMock]:
|
||||
"""Mock stream response."""
|
||||
|
||||
async def mock_generator(events: Iterable[RawMessageStreamEvent], **kwargs):
|
||||
async def mock_generator(
|
||||
events: Iterable[RawMessageStreamEvent],
|
||||
**kwargs: Unpack[MessageCreateParamsStreaming],
|
||||
) -> AsyncGenerator[RawMessageStreamEvent]:
|
||||
"""Create a stream of messages with the specified content blocks."""
|
||||
stop_reason = "end_turn"
|
||||
container = None
|
||||
has_message_delta = False
|
||||
refusal_magic_string = (
|
||||
"ANTHROPIC_MAGIC_STRING_TRIGGER_REFUSAL_"
|
||||
"1FAEFB6177B4672DEE07F9D3AFC62588"
|
||||
"CCD2631EDCF22E8CCC1FB35B501C9C86"
|
||||
)
|
||||
for message in kwargs.get("messages"):
|
||||
for message in kwargs["messages"]:
|
||||
if message["role"] != "user":
|
||||
continue
|
||||
if isinstance(message["content"], str):
|
||||
@@ -156,6 +162,8 @@ def mock_create_stream() -> Generator[AsyncMock]:
|
||||
type="message_start",
|
||||
)
|
||||
for event in events:
|
||||
if isinstance(event, RawMessageDeltaEvent):
|
||||
has_message_delta = True
|
||||
if isinstance(event, RawContentBlockStartEvent) and isinstance(
|
||||
event.content_block, ToolUseBlock
|
||||
):
|
||||
@@ -171,20 +179,21 @@ def mock_create_stream() -> Generator[AsyncMock]:
|
||||
]
|
||||
):
|
||||
container = Container(
|
||||
id=kwargs.get("container_id", "container_1234567890ABCDEFGHIJKLMN"),
|
||||
id=kwargs.get("container") or "container_1234567890ABCDEFGHIJKLMN",
|
||||
expires_at=dt_util.utcnow() + datetime.timedelta(minutes=5),
|
||||
)
|
||||
|
||||
yield event
|
||||
yield RawMessageDeltaEvent(
|
||||
type="message_delta",
|
||||
delta=Delta(
|
||||
stop_reason=stop_reason,
|
||||
stop_sequence="",
|
||||
container=container,
|
||||
),
|
||||
usage=MessageDeltaUsage(output_tokens=0),
|
||||
)
|
||||
if not has_message_delta:
|
||||
yield RawMessageDeltaEvent(
|
||||
type="message_delta",
|
||||
delta=Delta(
|
||||
stop_reason=stop_reason,
|
||||
stop_sequence="",
|
||||
container=container,
|
||||
),
|
||||
usage=MessageDeltaUsage(output_tokens=0),
|
||||
)
|
||||
yield RawMessageStopEvent(type="message_stop")
|
||||
|
||||
with patch(
|
||||
|
||||
@@ -1088,6 +1088,414 @@
|
||||
}),
|
||||
])
|
||||
# ---
|
||||
# name: test_resume_pause_turn[complete_on_last_request]
|
||||
list([
|
||||
dict({
|
||||
'content': 'Calculate the answer',
|
||||
'role': 'user',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_0',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_0',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_1',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_1',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_2',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_2',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_3',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_3',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_4',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_4',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_5',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_5',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_6',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_6',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_7',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_7',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_8',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
])
|
||||
# ---
|
||||
# name: test_resume_pause_turn[multiple_pauses]
|
||||
list([
|
||||
dict({
|
||||
'content': 'Calculate the answer',
|
||||
'role': 'user',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_0',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_0',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_1',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'content': dict({
|
||||
'content': list([
|
||||
]),
|
||||
'return_code': 0,
|
||||
'stderr': '',
|
||||
'stdout': '''
|
||||
42
|
||||
|
||||
''',
|
||||
'type': 'bash_code_execution_result',
|
||||
}),
|
||||
'tool_use_id': 'srvtoolu_1',
|
||||
'type': 'bash_code_execution_tool_result',
|
||||
}),
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_2',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
])
|
||||
# ---
|
||||
# name: test_resume_pause_turn[one_pause]
|
||||
list([
|
||||
dict({
|
||||
'content': 'Calculate the answer',
|
||||
'role': 'user',
|
||||
}),
|
||||
dict({
|
||||
'content': list([
|
||||
dict({
|
||||
'signature': 'ErUBCkYIARgCIkCYXaVNJShe3A86Hp7XUzh9YsCYBbJTbQsrklTAPtJ2sP/NoB6tSzpK/nTL6CjSo2R6n0KNBIg5MH6asM2R/kmaEgyB/X1FtZq5OQAC7jUaDEPWCdcwGQ4RaBy5wiIwmRxExIlDhoY6tILoVPnOExkC/0igZxHEwxK8RU/fmw0b+o+TwAarzUitwzbo21E5Kh3pa3I6yqVROf1t2F8rFocNUeCegsWV/ytwYV+ayA==',
|
||||
'thinking': 'I will calculate the answer.',
|
||||
'type': 'thinking',
|
||||
}),
|
||||
dict({
|
||||
'id': 'srvtoolu_0',
|
||||
'input': dict({
|
||||
'command': 'echo 42',
|
||||
}),
|
||||
'name': 'bash_code_execution',
|
||||
'type': 'server_tool_use',
|
||||
}),
|
||||
]),
|
||||
'role': 'assistant',
|
||||
}),
|
||||
])
|
||||
# ---
|
||||
# name: test_text_editor_code_execution[create_file]
|
||||
list([
|
||||
dict({
|
||||
|
||||
@@ -1,10 +1,19 @@
|
||||
"""Tests for the Anthropic integration."""
|
||||
|
||||
from json import JSONDecodeError
|
||||
from pathlib import Path
|
||||
import re
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from anthropic.types import Message, TextBlock, Usage
|
||||
from anthropic.types import (
|
||||
Message,
|
||||
MessageDeltaUsage,
|
||||
RawMessageDeltaEvent,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
Usage,
|
||||
)
|
||||
from anthropic.types.raw_message_delta_event import Delta
|
||||
from freezegun import freeze_time
|
||||
import probatio
|
||||
import pytest
|
||||
@@ -15,7 +24,7 @@ from homeassistant.core import HomeAssistant
|
||||
from homeassistant.exceptions import HomeAssistantError
|
||||
from homeassistant.helpers import entity_registry as er, selector
|
||||
|
||||
from . import create_content_block
|
||||
from . import create_content_block, create_server_tool_use_block
|
||||
|
||||
from tests.common import MockConfigEntry
|
||||
|
||||
@@ -108,6 +117,121 @@ async def test_stream_wrong_type(
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_init_component")
|
||||
@pytest.mark.parametrize(
|
||||
("stop_reason", "translation_key", "message"),
|
||||
[
|
||||
pytest.param(
|
||||
"max_tokens",
|
||||
"response_max_tokens",
|
||||
"Claude reached the output token limit before completing the response",
|
||||
id="max_tokens",
|
||||
),
|
||||
pytest.param(
|
||||
"model_context_window_exceeded",
|
||||
"response_context_window_exceeded",
|
||||
"Claude reached the context window limit before completing the response",
|
||||
id="context_window_exceeded",
|
||||
),
|
||||
pytest.param(
|
||||
"refusal",
|
||||
"api_refusal",
|
||||
"Potential policy violation detected",
|
||||
id="refusal",
|
||||
),
|
||||
pytest.param(
|
||||
"stop_sequence",
|
||||
"response_stop_sequence",
|
||||
"Claude stopped after encountering a stop sequence",
|
||||
id="stop_sequence",
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
("structure", "response_text"),
|
||||
[
|
||||
pytest.param(None, "The generated data starts with", id="plain"),
|
||||
pytest.param(
|
||||
probatio.Schema(
|
||||
{
|
||||
probatio.Required("characters"): selector.selector(
|
||||
{"text": {"multiple": True}}
|
||||
)
|
||||
}
|
||||
),
|
||||
'{"characters": ["Mario',
|
||||
id="structured",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_generate_data_stop_reason_error(
|
||||
hass: HomeAssistant,
|
||||
mock_create_stream: AsyncMock,
|
||||
stop_reason: StopReason,
|
||||
translation_key: str,
|
||||
message: str,
|
||||
structure: probatio.Schema | None,
|
||||
response_text: str,
|
||||
) -> None:
|
||||
"""Reject unsuccessful stop reasons before parsing structured data."""
|
||||
mock_create_stream.return_value = [
|
||||
[
|
||||
*create_content_block(0, [response_text]),
|
||||
RawMessageDeltaEvent(
|
||||
type="message_delta",
|
||||
delta=Delta(stop_reason=stop_reason),
|
||||
usage=MessageDeltaUsage(output_tokens=10),
|
||||
),
|
||||
]
|
||||
]
|
||||
|
||||
with pytest.raises(HomeAssistantError, match=re.escape(message)) as exc_info:
|
||||
await ai_task.async_generate_data(
|
||||
hass,
|
||||
task_name="Test Task",
|
||||
entity_id="ai_task.claude_ai_task",
|
||||
instructions="Generate test data",
|
||||
structure=structure,
|
||||
)
|
||||
|
||||
assert exc_info.value.translation_key == translation_key
|
||||
mock_create_stream.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_init_component")
|
||||
async def test_generate_data_invalid_tool_arguments(
|
||||
hass: HomeAssistant,
|
||||
mock_create_stream: AsyncMock,
|
||||
) -> None:
|
||||
"""Report invalid tool arguments with the first parsing error as the cause."""
|
||||
incomplete_json = '{"command": "echo'
|
||||
mock_create_stream.return_value = [
|
||||
[
|
||||
*create_server_tool_use_block(
|
||||
0, "srvtoolu_first", "bash_code_execution", [incomplete_json]
|
||||
),
|
||||
*create_server_tool_use_block(
|
||||
1, "srvtoolu_second", "bash_code_execution", ['{"command":']
|
||||
),
|
||||
]
|
||||
]
|
||||
|
||||
with pytest.raises(
|
||||
HomeAssistantError, match="Claude returned invalid tool arguments"
|
||||
) as exc_info:
|
||||
await ai_task.async_generate_data(
|
||||
hass,
|
||||
task_name="Test Task",
|
||||
entity_id="ai_task.claude_ai_task",
|
||||
instructions="Generate test data",
|
||||
)
|
||||
|
||||
assert exc_info.value.translation_key == "tool_args_parse_error"
|
||||
assert isinstance(exc_info.value.__cause__, JSONDecodeError)
|
||||
assert exc_info.value.__cause__.doc == incomplete_json
|
||||
mock_create_stream.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_init_component")
|
||||
async def test_generate_invalid_structured_data(
|
||||
hass: HomeAssistant,
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
"""Tests for the Anthropic integration."""
|
||||
|
||||
from collections.abc import AsyncIterator, Generator
|
||||
from copy import deepcopy
|
||||
import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Unpack
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
from anthropic import RateLimitError
|
||||
@@ -12,6 +14,7 @@ from anthropic.types import (
|
||||
CitationsConfig,
|
||||
CitationsWebSearchResultLocation,
|
||||
CitationWebSearchResultLocationParam,
|
||||
Container,
|
||||
DocumentBlock,
|
||||
EncryptedCodeExecutionResultBlock,
|
||||
Message,
|
||||
@@ -20,7 +23,9 @@ from anthropic.types import (
|
||||
RawMessageDeltaEvent,
|
||||
RawMessageStartEvent,
|
||||
RawMessageStopEvent,
|
||||
RawMessageStreamEvent,
|
||||
ServerToolCaller20260120,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
TextEditorCodeExecutionCreateResultBlock,
|
||||
TextEditorCodeExecutionStrReplaceResultBlock,
|
||||
@@ -34,6 +39,7 @@ from anthropic.types import (
|
||||
WebSearchResultBlock,
|
||||
WebSearchToolResultError,
|
||||
)
|
||||
from anthropic.types.message_create_params import MessageCreateParamsStreaming
|
||||
from anthropic.types.raw_message_delta_event import Delta
|
||||
from anthropic.types.text_editor_code_execution_tool_result_block import (
|
||||
Content as TextEditorCodeExecutionToolResultBlockContent,
|
||||
@@ -64,6 +70,7 @@ from homeassistant.components.anthropic.const import (
|
||||
DOMAIN,
|
||||
)
|
||||
from homeassistant.components.anthropic.entity import (
|
||||
MAX_TOOL_ITERATIONS,
|
||||
CitationDetails,
|
||||
ContentDetails,
|
||||
_convert_content,
|
||||
@@ -83,7 +90,7 @@ from homeassistant.helpers import (
|
||||
llm,
|
||||
)
|
||||
from homeassistant.setup import async_setup_component
|
||||
from homeassistant.util import ulid as ulid_util
|
||||
from homeassistant.util import dt as dt_util, ulid as ulid_util
|
||||
|
||||
from . import (
|
||||
create_bash_code_execution_result_block,
|
||||
@@ -101,6 +108,68 @@ from . import (
|
||||
|
||||
from tests.common import MockConfigEntry
|
||||
|
||||
ENTITY_ID = "conversation.claude_conversation"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_config_entry_with_server_tools(
|
||||
hass: HomeAssistant, mock_config_entry: MockConfigEntry
|
||||
) -> MockConfigEntry:
|
||||
"""Configure server tools and adaptive thinking."""
|
||||
hass.config_entries.async_update_subentry(
|
||||
mock_config_entry,
|
||||
next(iter(mock_config_entry.subentries.values())),
|
||||
data={
|
||||
CONF_CHAT_MODEL: "claude-opus-4-7",
|
||||
CONF_CODE_EXECUTION: True,
|
||||
CONF_THINKING_EFFORT: "medium",
|
||||
CONF_LLM_HASS_API: llm.LLM_API_ASSIST,
|
||||
},
|
||||
)
|
||||
return mock_config_entry
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def captured_requests(
|
||||
mock_create_stream: AsyncMock,
|
||||
) -> list[MessageCreateParamsStreaming]:
|
||||
"""Capture request content before the conversation mutates it."""
|
||||
requests: list[MessageCreateParamsStreaming] = []
|
||||
create_stream = mock_create_stream.side_effect
|
||||
|
||||
def capture_request(
|
||||
**kwargs: Unpack[MessageCreateParamsStreaming],
|
||||
) -> AsyncIterator[RawMessageStreamEvent]:
|
||||
requests.append(deepcopy(kwargs))
|
||||
return create_stream(**kwargs)
|
||||
|
||||
mock_create_stream.side_effect = capture_request
|
||||
return requests
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def code_execution_container() -> Container:
|
||||
"""Return an active code execution container."""
|
||||
return Container(
|
||||
id="container_paused",
|
||||
expires_at=dt_util.utcnow() + datetime.timedelta(minutes=5),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_llm_tool() -> Generator[AsyncMock]:
|
||||
"""Provide a local tool whose execution can be checked."""
|
||||
mock_tool = AsyncMock()
|
||||
mock_tool.name = "test_tool"
|
||||
mock_tool.description = "Test function"
|
||||
mock_tool.parameters = probatio.Schema({probatio.Optional("param1"): str})
|
||||
mock_tool.async_call.return_value = llm.ToolResult(data="Test response")
|
||||
with patch(
|
||||
"homeassistant.components.llm.async_get_tools",
|
||||
return_value=LLMTools(tools=[mock_tool]),
|
||||
):
|
||||
yield mock_tool
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_init_component")
|
||||
async def test_entity(
|
||||
@@ -694,15 +763,47 @@ async def test_conversation_id(
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_init_component")
|
||||
async def test_refusal(
|
||||
@pytest.mark.parametrize(
|
||||
("stop_reason", "error_message"),
|
||||
[
|
||||
pytest.param(
|
||||
"refusal",
|
||||
"Potential policy violation detected",
|
||||
id="refusal",
|
||||
),
|
||||
pytest.param(
|
||||
"max_tokens",
|
||||
"Claude reached the output token limit before completing the response",
|
||||
id="output_token_limit",
|
||||
),
|
||||
pytest.param(
|
||||
"model_context_window_exceeded",
|
||||
"Claude reached the context window limit before completing the response",
|
||||
id="context_window_limit",
|
||||
),
|
||||
pytest.param(
|
||||
"stop_sequence",
|
||||
"Claude stopped after encountering a stop sequence",
|
||||
id="stop_sequence",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_stop_reason_error(
|
||||
hass: HomeAssistant,
|
||||
mock_create_stream: AsyncMock,
|
||||
stop_reason: StopReason,
|
||||
error_message: str,
|
||||
) -> None:
|
||||
"""Test refusal due to potential policy violation."""
|
||||
"""Test errors for refused, truncated, or stop-sequence responses."""
|
||||
mock_create_stream.return_value = [
|
||||
create_content_block(
|
||||
0, ["Certainly! To take over the world you need just a simple "]
|
||||
)
|
||||
[
|
||||
*create_content_block(0, ["An incomplete response"]),
|
||||
RawMessageDeltaEvent(
|
||||
type="message_delta",
|
||||
delta=Delta(stop_reason=stop_reason),
|
||||
usage=MessageDeltaUsage(output_tokens=10),
|
||||
),
|
||||
]
|
||||
]
|
||||
|
||||
result = await conversation.async_converse(
|
||||
@@ -711,16 +812,124 @@ async def test_refusal(
|
||||
"EDCF22E8CCC1FB35B501C9C86",
|
||||
None,
|
||||
Context(),
|
||||
agent_id="conversation.claude_conversation",
|
||||
agent_id=ENTITY_ID,
|
||||
)
|
||||
|
||||
assert result.response.response_type is intent.IntentResponseType.ERROR
|
||||
assert result.response.error_code == "unknown"
|
||||
assert (
|
||||
result.response.speech["plain"]["speech"]
|
||||
== "Potential policy violation detected"
|
||||
assert result.response.speech["plain"]["speech"] == error_message
|
||||
mock_create_stream.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_config_entry_with_server_tools", "mock_init_component")
|
||||
@pytest.mark.parametrize(
|
||||
("tool_blocks", "stop_reason", "error_message"),
|
||||
[
|
||||
pytest.param(
|
||||
create_tool_use_block(0, "toolu_invalid", "test_tool", ['{"param1":']),
|
||||
"max_tokens",
|
||||
"Claude reached the output token limit before completing the response",
|
||||
id="local_output_token_limit",
|
||||
),
|
||||
pytest.param(
|
||||
create_server_tool_use_block(
|
||||
0, "srvtoolu_invalid", "bash_code_execution", ['{"command":']
|
||||
),
|
||||
"max_tokens",
|
||||
"Claude reached the output token limit before completing the response",
|
||||
id="server_output_token_limit",
|
||||
),
|
||||
pytest.param(
|
||||
create_tool_use_block(0, "toolu_invalid", "test_tool", ['{"param1":']),
|
||||
"model_context_window_exceeded",
|
||||
"Claude reached the context window limit before completing the response",
|
||||
id="context_window_limit",
|
||||
),
|
||||
pytest.param(
|
||||
create_tool_use_block(0, "toolu_invalid", "test_tool", ['{"param1":']),
|
||||
"refusal",
|
||||
"Potential policy violation detected",
|
||||
id="refusal",
|
||||
),
|
||||
pytest.param(
|
||||
create_tool_use_block(0, "toolu_invalid", "test_tool", ['{"param1":']),
|
||||
"stop_sequence",
|
||||
"Claude stopped after encountering a stop sequence",
|
||||
id="stop_sequence",
|
||||
),
|
||||
pytest.param(
|
||||
create_tool_use_block(0, "toolu_invalid", "test_tool", ['{"param1":']),
|
||||
"tool_use",
|
||||
"Claude returned invalid tool arguments",
|
||||
id="local_invalid_json",
|
||||
),
|
||||
pytest.param(
|
||||
create_server_tool_use_block(
|
||||
0, "srvtoolu_invalid", "bash_code_execution", ['{"command":']
|
||||
),
|
||||
"end_turn",
|
||||
"Claude returned invalid tool arguments",
|
||||
id="server_invalid_json",
|
||||
),
|
||||
pytest.param(
|
||||
create_server_tool_use_block(
|
||||
0, "srvtoolu_invalid", "bash_code_execution", ['{"command":']
|
||||
),
|
||||
"pause_turn",
|
||||
"Claude returned invalid tool arguments",
|
||||
id="invalid_json_prevents_continuation",
|
||||
),
|
||||
pytest.param(
|
||||
[
|
||||
*create_tool_use_block(0, "toolu_invalid", "test_tool", ['{"param1":']),
|
||||
*create_tool_use_block(1, "toolu_valid", "test_tool", ["{}"]),
|
||||
],
|
||||
"tool_use",
|
||||
"Claude returned invalid tool arguments",
|
||||
id="valid_tool_after_invalid_tool",
|
||||
),
|
||||
pytest.param(
|
||||
[
|
||||
*create_tool_use_block(0, "toolu_invalid", "test_tool", ['{"param1":']),
|
||||
*create_tool_use_block(1, "toolu_invalid_2", "test_tool", ["{"]),
|
||||
*create_tool_use_block(2, "toolu_valid", "test_tool", ["{}"]),
|
||||
],
|
||||
"tool_use",
|
||||
"Claude returned invalid tool arguments",
|
||||
id="valid_tool_after_multiple_invalid_tools",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_invalid_tool_arguments(
|
||||
hass: HomeAssistant,
|
||||
mock_create_stream: AsyncMock,
|
||||
mock_llm_tool: AsyncMock,
|
||||
tool_blocks: list[RawMessageStreamEvent],
|
||||
stop_reason: StopReason,
|
||||
error_message: str,
|
||||
) -> None:
|
||||
"""Read the stop reason before reporting malformed tool arguments."""
|
||||
mock_create_stream.return_value = [
|
||||
[
|
||||
*tool_blocks,
|
||||
RawMessageDeltaEvent(
|
||||
type="message_delta",
|
||||
delta=Delta(stop_reason=stop_reason),
|
||||
usage=MessageDeltaUsage(output_tokens=10),
|
||||
),
|
||||
]
|
||||
]
|
||||
|
||||
result = await conversation.async_converse(
|
||||
hass, "Please call the test function", None, Context(), agent_id=ENTITY_ID
|
||||
)
|
||||
|
||||
assert result.response.response_type is intent.IntentResponseType.ERROR
|
||||
assert result.response.error_code == "unknown"
|
||||
assert result.response.speech["plain"]["speech"] == error_message
|
||||
mock_llm_tool.async_call.assert_not_called()
|
||||
mock_create_stream.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_init_component")
|
||||
async def test_stream_wrong_type(
|
||||
@@ -2110,6 +2319,185 @@ async def test_container_reused(
|
||||
assert mock_create_stream.call_args.kwargs["container"] == container_id
|
||||
|
||||
|
||||
def _create_paused_response(
|
||||
tool_number: int, container: Container, index: int = 0
|
||||
) -> list[RawMessageStreamEvent]:
|
||||
"""Create a response ending with a pending server tool call."""
|
||||
return [
|
||||
*create_thinking_block(index, ["I will calculate the answer."]),
|
||||
*create_server_tool_use_block(
|
||||
index + 1,
|
||||
f"srvtoolu_{tool_number}",
|
||||
"bash_code_execution",
|
||||
['{"command": "echo 42"}'],
|
||||
),
|
||||
RawMessageDeltaEvent(
|
||||
type="message_delta",
|
||||
delta=Delta(stop_reason="pause_turn", container=container),
|
||||
usage=MessageDeltaUsage(output_tokens=10),
|
||||
),
|
||||
# Usage updates must not erase the preceding stop reason or container.
|
||||
RawMessageDeltaEvent(
|
||||
type="message_delta",
|
||||
delta=Delta(),
|
||||
usage=MessageDeltaUsage(output_tokens=11),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _create_paused_responses(
|
||||
pause_count: int, container: Container
|
||||
) -> list[list[RawMessageStreamEvent]]:
|
||||
"""Complete each preceding server tool before pausing on another one."""
|
||||
return [
|
||||
_create_paused_response(0, container),
|
||||
*[
|
||||
[
|
||||
*create_bash_code_execution_result_block(
|
||||
0, f"srvtoolu_{tool_number - 1}", stdout="42\n"
|
||||
),
|
||||
*_create_paused_response(tool_number, container, index=1),
|
||||
]
|
||||
for tool_number in range(1, pause_count)
|
||||
],
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_config_entry_with_server_tools", "mock_init_component")
|
||||
@pytest.mark.parametrize(
|
||||
"pause_count",
|
||||
[
|
||||
pytest.param(1, id="one_pause"),
|
||||
pytest.param(3, id="multiple_pauses"),
|
||||
pytest.param(MAX_TOOL_ITERATIONS - 1, id="complete_on_last_request"),
|
||||
],
|
||||
)
|
||||
async def test_resume_pause_turn(
|
||||
hass: HomeAssistant,
|
||||
mock_create_stream: AsyncMock,
|
||||
captured_requests: list[MessageCreateParamsStreaming],
|
||||
code_execution_container: Container,
|
||||
snapshot: SnapshotAssertion,
|
||||
pause_count: int,
|
||||
) -> None:
|
||||
"""Resume pending server tools with the existing response and configuration."""
|
||||
mock_create_stream.return_value = [
|
||||
*_create_paused_responses(pause_count, code_execution_container),
|
||||
[
|
||||
*create_bash_code_execution_result_block(
|
||||
0, f"srvtoolu_{pause_count - 1}", stdout="42\n"
|
||||
),
|
||||
*create_content_block(1, ["The answer is 42."]),
|
||||
],
|
||||
]
|
||||
|
||||
result = await conversation.async_converse(
|
||||
hass, "Calculate the answer", None, Context(), agent_id=ENTITY_ID
|
||||
)
|
||||
|
||||
assert result.response.response_type is intent.IntentResponseType.ACTION_DONE
|
||||
assert result.response.speech["plain"]["speech"] == "The answer is 42."
|
||||
assert mock_create_stream.await_count == pause_count + 1
|
||||
assert captured_requests[-1]["messages"] == snapshot
|
||||
for request in captured_requests[1:]:
|
||||
assert request["container"] == code_execution_container.id
|
||||
assert request["tools"] == captured_requests[0]["tools"]
|
||||
assert request["thinking"] == captured_requests[0]["thinking"]
|
||||
assert request["output_config"] == captured_requests[0]["output_config"]
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_config_entry_with_server_tools", "mock_init_component")
|
||||
@pytest.mark.parametrize(
|
||||
"final_tool_arguments",
|
||||
[
|
||||
pytest.param(['{"command": "echo 42"}'], id="valid_tool_arguments"),
|
||||
pytest.param(['{"command":'], id="invalid_tool_arguments"),
|
||||
],
|
||||
)
|
||||
async def test_pause_turn_iteration_limit(
|
||||
hass: HomeAssistant,
|
||||
mock_create_stream: AsyncMock,
|
||||
code_execution_container: Container,
|
||||
final_tool_arguments: list[str],
|
||||
) -> None:
|
||||
"""Report an incomplete response if every allowed request pauses."""
|
||||
mock_create_stream.return_value = [
|
||||
*_create_paused_responses(MAX_TOOL_ITERATIONS - 1, code_execution_container),
|
||||
[
|
||||
*create_bash_code_execution_result_block(
|
||||
0, f"srvtoolu_{MAX_TOOL_ITERATIONS - 2}", stdout="42\n"
|
||||
),
|
||||
*create_server_tool_use_block(
|
||||
1, "srvtoolu_final", "bash_code_execution", final_tool_arguments
|
||||
),
|
||||
RawMessageDeltaEvent(
|
||||
type="message_delta",
|
||||
delta=Delta(
|
||||
stop_reason="pause_turn", container=code_execution_container
|
||||
),
|
||||
usage=MessageDeltaUsage(output_tokens=10),
|
||||
),
|
||||
],
|
||||
]
|
||||
|
||||
result = await conversation.async_converse(
|
||||
hass, "Calculate the answer", None, Context(), agent_id=ENTITY_ID
|
||||
)
|
||||
|
||||
assert mock_create_stream.await_count == MAX_TOOL_ITERATIONS
|
||||
assert result.response.response_type is intent.IntentResponseType.ERROR
|
||||
assert result.response.error_code == "unknown"
|
||||
assert result.response.speech["plain"]["speech"] == (
|
||||
"Claude could not complete the response within the allowed number of requests"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("mock_config_entry_with_server_tools", "mock_init_component")
|
||||
async def test_pause_turn_followed_by_local_tool(
|
||||
hass: HomeAssistant,
|
||||
mock_create_stream: AsyncMock,
|
||||
captured_requests: list[MessageCreateParamsStreaming],
|
||||
code_execution_container: Container,
|
||||
) -> None:
|
||||
"""Continue processing local tools after resuming a paused server tool."""
|
||||
mock_tool = AsyncMock()
|
||||
mock_tool.name = "test_tool"
|
||||
mock_tool.description = "Test function"
|
||||
mock_tool.parameters = probatio.Schema({})
|
||||
mock_tool.async_call.return_value = llm.ToolResult(data="Test response")
|
||||
mock_create_stream.return_value = [
|
||||
_create_paused_response(0, code_execution_container),
|
||||
[
|
||||
*create_bash_code_execution_result_block(0, "srvtoolu_0", stdout="42\n"),
|
||||
*create_tool_use_block(1, "toolu_local", "test_tool", ["{}"]),
|
||||
],
|
||||
create_content_block(0, ["The answer is 42."]),
|
||||
]
|
||||
|
||||
with patch(
|
||||
"homeassistant.components.llm.async_get_tools",
|
||||
return_value=LLMTools(tools=[mock_tool]),
|
||||
):
|
||||
result = await conversation.async_converse(
|
||||
hass, "Calculate the answer", None, Context(), agent_id=ENTITY_ID
|
||||
)
|
||||
|
||||
assert result.response.speech["plain"]["speech"] == "The answer is 42."
|
||||
assert mock_create_stream.await_count == 3
|
||||
mock_tool.async_call.assert_awaited_once()
|
||||
assert list(captured_requests[2]["messages"])[-1] == {
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_local",
|
||||
"content": '"Test response"',
|
||||
"is_error": False,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"content",
|
||||
[
|
||||
|
||||
Reference in New Issue
Block a user