diff --git a/homeassistant/components/anthropic/entity.py b/homeassistant/components/anthropic/entity.py index 4dacf98bb920..8b0fd76f5280 100644 --- a/homeassistant/components/anthropic/entity.py +++ b/homeassistant/components/anthropic/entity.py @@ -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 diff --git a/homeassistant/components/anthropic/strings.json b/homeassistant/components/anthropic/strings.json index 012f6ef2f2e2..3a2426812c8b 100644 --- a/homeassistant/components/anthropic/strings.json +++ b/homeassistant/components/anthropic/strings.json @@ -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}." }, diff --git a/tests/components/anthropic/conftest.py b/tests/components/anthropic/conftest.py index 0909aac2da6d..f65a5bcb1419 100644 --- a/tests/components/anthropic/conftest.py +++ b/tests/components/anthropic/conftest.py @@ -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( diff --git a/tests/components/anthropic/snapshots/test_conversation.ambr b/tests/components/anthropic/snapshots/test_conversation.ambr index a927d418fb91..ba85af8548eb 100644 --- a/tests/components/anthropic/snapshots/test_conversation.ambr +++ b/tests/components/anthropic/snapshots/test_conversation.ambr @@ -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({ diff --git a/tests/components/anthropic/test_ai_task.py b/tests/components/anthropic/test_ai_task.py index 75fb112e6cb3..f06f59233102 100644 --- a/tests/components/anthropic/test_ai_task.py +++ b/tests/components/anthropic/test_ai_task.py @@ -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, diff --git a/tests/components/anthropic/test_conversation.py b/tests/components/anthropic/test_conversation.py index 74981d6ecbca..b5a70b485266 100644 --- a/tests/components/anthropic/test_conversation.py +++ b/tests/components/anthropic/test_conversation.py @@ -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", [