Close Anthropic connection on service cancellation (#183640)

This commit is contained in:
Denis Shulyaka
2026-09-29 22:57:27 +02:00
committed by GitHub
parent 7bccc2eb4e
commit f2d85f409f
3 changed files with 39 additions and 21 deletions
+13 -13
View File
@@ -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
+13 -2
View File
@@ -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
)
@@ -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",