Handle Anthropic stop reasons (#183384)

This commit is contained in:
Denis Shulyaka
2026-09-29 18:37:03 +02:00
committed by GitHub
parent 95fb45a7f4
commit 9dccac7fa5
6 changed files with 1035 additions and 51 deletions
+66 -26
View File
@@ -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}."
},
+21 -12
View File
@@ -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({
+126 -2
View File
@@ -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,
+399 -11
View File
@@ -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",
[