Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 18 additions & 2 deletions src/anthropic/lib/bedrock/_stream_decoder.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING, Iterator, AsyncIterator

from ..._utils import lru_cache
Expand Down Expand Up @@ -37,7 +38,9 @@ def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[ServerSentEvent]:
for event in event_stream_buffer:
message = self._parse_message_from_event(event)
if message:
yield ServerSentEvent(data=message, event="completion")
sse = self._build_sse(message)
if sse is not None:
yield sse

async def aiter_bytes(self, iterator: AsyncIterator[bytes]) -> AsyncIterator[ServerSentEvent]:
"""Given an async iterator that yields lines, iterate over it & yield every event encountered"""
Expand All @@ -49,7 +52,20 @@ async def aiter_bytes(self, iterator: AsyncIterator[bytes]) -> AsyncIterator[Ser
for event in event_stream_buffer:
message = self._parse_message_from_event(event)
if message:
yield ServerSentEvent(data=message, event="completion")
sse = self._build_sse(message)
if sse is not None:
yield sse

def _build_sse(self, message: str) -> ServerSentEvent | None:
payload = json.loads(message)
if not isinstance(payload, dict):
return None

event_type = payload.get("type")
if not isinstance(event_type, str):
return None

return ServerSentEvent(data=message, event=event_type)

def _parse_message_from_event(self, event: EventStreamMessage) -> str | None:
response_dict = event.to_response_dict()
Expand Down
19 changes: 19 additions & 0 deletions tests/lib/test_bedrock_stream_decoder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
from anthropic.lib.bedrock._stream_decoder import AWSEventStreamDecoder


def test_build_sse_uses_payload_type() -> None:
decoder = AWSEventStreamDecoder()

sse = decoder._build_sse('{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}')

assert sse is not None
assert sse.event == "content_block_delta"
assert sse.json() == {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hi"}}


def test_build_sse_drops_non_message_payloads() -> None:
decoder = AWSEventStreamDecoder()

sse = decoder._build_sse('{"amazon-bedrock-invocationMetrics":{"inputTokenCount":1}}')

assert sse is None