From f2d85f409fb0dd9a60af7701739dee1fa62b2763 Mon Sep 17 00:00:00 2001 From: Denis Shulyaka Date: Tue, 29 Sep 2026 23:57:27 +0300 Subject: [PATCH] Close Anthropic connection on service cancellation (#183640) --- homeassistant/components/anthropic/entity.py | 26 +++++++++---------- tests/components/anthropic/conftest.py | 15 +++++++++-- .../components/anthropic/test_conversation.py | 19 +++++++++----- 3 files changed, 39 insertions(+), 21 deletions(-) diff --git a/homeassistant/components/anthropic/entity.py b/homeassistant/components/anthropic/entity.py index 8b0fd76f5280..91958af0df09 100644 --- a/homeassistant/components/anthropic/entity.py +++ b/homeassistant/components/anthropic/entity.py @@ -528,6 +528,10 @@ class AnthropicDeltaStream: stream: AsyncStream[MessageStreamEvent], ) -> None: """Initialize the delta stream.""" + if not hasattr(stream, "__aiter__"): + raise HomeAssistantError( + translation_domain=DOMAIN, translation_key="unexpected_stream_object" + ) self._chat_log: conversation.ChatLog = chat_log self._stream: AsyncStream[MessageStreamEvent] = stream self.stop_reason: StopReason | None = None @@ -554,10 +558,6 @@ class AnthropicDeltaStream: conversation.AssistantContentDeltaDict | conversation.ToolResultContentDeltaDict ]: """Initialize the stream and return the async iterator.""" - if self._stream is None or not hasattr(self._stream, "__aiter__"): - raise HomeAssistantError( - translation_domain=DOMAIN, translation_key="unexpected_stream_object" - ) if self._stream_iterator is None: self._stream_iterator = self._stream.__aiter__() return self @@ -1145,15 +1145,15 @@ class AnthropicBaseLLMEntity(CoordinatorEntity[AnthropicCoordinator]): 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, - delta_stream, - ) - ] - ) + async with stream: + new_messages, model_args["container"] = _convert_content( + [ + content + async for content in chat_log.async_add_delta_content_stream( + self.entity_id, delta_stream + ) + ] + ) cast(list[MessageParam], model_args["messages"]).extend(new_messages) except anthropic.AuthenticationError as err: # Trigger coordinator to confirm the auth failure diff --git a/tests/components/anthropic/conftest.py b/tests/components/anthropic/conftest.py index f65a5bcb1419..4c98a7b07166 100644 --- a/tests/components/anthropic/conftest.py +++ b/tests/components/anthropic/conftest.py @@ -3,8 +3,9 @@ from collections.abc import AsyncGenerator, Generator, Iterable import datetime from typing import Unpack -from unittest.mock import DEFAULT, AsyncMock, patch +from unittest.mock import DEFAULT, AsyncMock, MagicMock, patch +from anthropic import AsyncStream from anthropic.pagination import AsyncPage from anthropic.types import ( Container, @@ -196,12 +197,22 @@ def mock_create_stream() -> Generator[AsyncMock]: ) yield RawMessageStopEvent(type="message_stop") + def mock_stream( + events: Iterable[RawMessageStreamEvent], + **kwargs: Unpack[MessageCreateParamsStreaming], + ) -> MagicMock: + """Create a stream supporting asynchronous iteration and cleanup.""" + stream = MagicMock(spec=AsyncStream) + stream.__aenter__.return_value = stream + stream.__aiter__.side_effect = lambda: mock_generator(events, **kwargs) + return stream + with patch( "anthropic.resources.messages.AsyncMessages.create", new_callable=AsyncMock, ) as mock_create: mock_create.side_effect = lambda **kwargs: ( - mock_generator(mock_create.return_value.pop(0), **kwargs) + mock_stream(mock_create.return_value.pop(0), **kwargs) if isinstance(mock_create.return_value, list) else DEFAULT ) diff --git a/tests/components/anthropic/test_conversation.py b/tests/components/anthropic/test_conversation.py index b5a70b485266..b109cae6b448 100644 --- a/tests/components/anthropic/test_conversation.py +++ b/tests/components/anthropic/test_conversation.py @@ -4,10 +4,10 @@ from collections.abc import AsyncIterator, Generator from copy import deepcopy import datetime from pathlib import Path -from typing import Any, Unpack -from unittest.mock import AsyncMock, Mock, patch +from typing import Unpack +from unittest.mock import AsyncMock, MagicMock, Mock, patch -from anthropic import RateLimitError +from anthropic import AsyncStream, RateLimitError from anthropic.types import ( CitationCharLocation, CitationCharLocationParam, @@ -343,7 +343,9 @@ async def test_token_stats_reported( """Test that cache reads, not cache creation, are reported as cached tokens.""" trace.async_clear_traces() - async def mock_stream(**kwargs: Any): + async def mock_stream( + **kwargs: Unpack[MessageCreateParamsStreaming], + ) -> AsyncIterator[RawMessageStreamEvent]: """Stream a single response carrying distinct cache read and creation usage.""" yield RawMessageStartEvent( type="message_start", @@ -370,11 +372,16 @@ async def test_token_stats_reported( ) yield RawMessageStopEvent(type="message_stop") + stream = MagicMock(spec=AsyncStream) + stream.__aenter__.return_value = stream with patch( "anthropic.resources.messages.AsyncMessages.create", new_callable=AsyncMock, - side_effect=mock_stream, - ): + return_value=stream, + ) as mock_create: + stream.__aiter__.side_effect = lambda: mock_stream( + **mock_create.call_args.kwargs + ) await conversation.async_converse( hass, "hello",