mirror of
https://github.com/home-assistant/core.git
synced 2026-10-06 14:29:21 -04:00
Close Anthropic connection on service cancellation (#183640)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user