This commit is contained in:
Kent 2026-09-15 20:27:11 -04:00 committed by GitHub
commit e0d66ba188
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 353 additions and 236 deletions

View file

@ -5,7 +5,8 @@ https://docs.aws.amazon.com/bedrock/latest/APIReference/API_agentcore_InvokeAgen
"""
import json
from collections.abc import AsyncGenerator
from collections.abc import AsyncGenerator, Iterator
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, Optional, Union, cast
from urllib.parse import quote
@ -51,6 +52,27 @@ else:
AsyncHTTPHandler = Any
_SSE_SENTINEL_PAYLOADS: Final = frozenset({"[DONE]"})
@dataclass(frozen=True, slots=True)
class _SseTextDelta:
text: str
@dataclass(frozen=True, slots=True)
class _SseUsage:
usage: AgentCoreUsage
@dataclass(frozen=True, slots=True)
class _SseFinalMessage:
message: AgentCoreMessage
_SseEvent = _SseTextDelta | _SseUsage | _SseFinalMessage
class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
def __init__(self, **kwargs):
BaseConfig.__init__(self, **kwargs)
@ -291,22 +313,95 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
return bool(value)
return False
def _extract_sse_json(self, line: str) -> dict | None:
"""Extract and parse JSON from an SSE data line."""
if not line.startswith("data:"):
def _extract_sse_payload(self, line: str) -> str | None:
"""Strip the 'data:' prefix from an SSE line, returning the payload or None."""
stripped: Final = line.strip()
if not stripped.startswith("data:"):
return None
payload: Final = stripped[5:].strip()
return payload or None
json_str: Final = line[5:].strip()
if not json_str:
return None
def _text_or_sentinel(self, text: str) -> tuple[_SseEvent, ...]:
"""Surface text as a content delta, dropping empty text and control sentinels."""
if not text or text in _SSE_SENTINEL_PAYLOADS:
return ()
return (_SseTextDelta(text),)
def _parse_sse_data(self, payload: str) -> tuple[_SseEvent, ...]:
"""Turn one SSE data payload into typed events.
Strands runtimes emit bare-string payloads (JSON-quoted like `"hi"` or
plain text like `hi`); these surface as content deltas instead of being
dropped. Control sentinels (e.g. `[DONE]`) are filtered whether quoted
or bare. A payload that fails to parse but looks like intended JSON
(starts with `{`, `[`, or `"`) is a malformed/truncated frame: it is
logged and dropped rather than leaked as content.
"""
try:
data: Final = json.loads(json_str)
# Skip non-dict data (some lines contain JSON strings)
return data if isinstance(data, dict) else None
data_obj: Final = json.loads(payload)
except json.JSONDecodeError:
verbose_logger.debug("Skipping non-JSON line: %s", line[:100])
return None
if payload not in _SSE_SENTINEL_PAYLOADS and payload[:1] in ("{", "[", '"'):
verbose_logger.debug("Skipping malformed JSON SSE line: %s", payload[:100])
return ()
return self._text_or_sentinel(payload)
if isinstance(data_obj, str):
return self._text_or_sentinel(data_obj)
if not isinstance(data_obj, dict):
return ()
return self._dict_to_sse_events(data_obj)
def _dict_to_sse_events(self, data_obj: dict[str, object]) -> tuple[_SseEvent, ...]:
event: Final = data_obj.get("event")
text: Final = self._extract_content_delta(data_obj) if isinstance(event, dict) else None
usage: Final = self._extract_usage_from_event(data_obj) if isinstance(event, dict) else None
message: Final = data_obj.get("message")
return (
*((_SseTextDelta(text),) if text else ()),
*((_SseUsage(usage),) if usage else ()),
*((_SseFinalMessage(cast(AgentCoreMessage, message)),) if isinstance(message, dict) else ()),
)
def _sse_event_to_chunk(self, event: _SseEvent, model: str) -> ModelResponseStream:
chunk: Final = ModelResponseStream(
id=f"chatcmpl-{uuid.uuid4()}",
created=0,
model=model,
object="chat.completion.chunk",
)
match event:
case _SseTextDelta(text=text):
chunk.choices = [
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content=text, role="assistant"),
)
]
case _SseUsage(usage=usage):
chunk.choices = [StreamingChoices(finish_reason="stop", index=0, delta=Delta())]
setattr(
chunk,
"usage",
Usage(
prompt_tokens=usage.get("inputTokens", 0),
completion_tokens=usage.get("outputTokens", 0),
total_tokens=usage.get("totalTokens", 0),
),
)
case _SseFinalMessage():
chunk.choices = [StreamingChoices(finish_reason="stop", index=0, delta=Delta())]
return chunk
def _line_to_chunks(self, line: str, model: str) -> Iterator[ModelResponseStream]:
"""Parse one raw SSE line and yield the resulting stream chunks."""
payload: Final = self._extract_sse_payload(line)
if payload is None:
return
for event in self._parse_sse_data(payload):
yield self._sse_event_to_chunk(event, model)
def _extract_usage_from_event(self, event_data: dict) -> AgentCoreUsage | None:
"""Extract usage information from event metadata."""
@ -480,51 +575,21 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
return self._parse_sse_stream(response_text)
def _parse_sse_stream(self, response_text: str) -> AgentCoreParsedResponse:
"""
Parse Server-Sent Events (SSE) stream format.
Each line starts with 'data:' followed by JSON.
"""Aggregate an SSE stream into a single parsed response."""
events: Final = tuple(
event
for line in response_text.strip().split("\n")
if (payload := self._extract_sse_payload(line)) is not None
for event in self._parse_sse_data(payload)
)
Returns:
AgentCoreParsedResponse: Parsed response with content, usage, and message
"""
final_message: AgentCoreMessage | None = None
usage_data: AgentCoreUsage | None = None
content_blocks: Final[list[str]] = []
for line in response_text.strip().split("\n"):
line = line.strip()
if not line:
continue
data = self._extract_sse_json(line)
if not data:
continue
verbose_logger.debug("SSE event keys: %s", list(data.keys()))
# Check for final complete message
if "message" in data and isinstance(data["message"], dict):
final_message = data["message"]
verbose_logger.debug("Found final message")
# Process event data
if "event" in data and isinstance(data["event"], dict):
event_payload = data["event"]
verbose_logger.debug("Event payload keys: %s", list(event_payload.keys()))
# Extract usage metadata
if usage := self._extract_usage_from_event(data):
usage_data = usage
verbose_logger.debug("Found usage data: %s", usage_data)
# Collect content deltas
if text := self._extract_content_delta(data):
content_blocks.append(text)
# Build final content
content: Final = self._extract_content_from_message(final_message) if final_message else "".join(content_blocks)
verbose_logger.debug("Final usage_data: %s", usage_data)
final_message: Final = next((e.message for e in reversed(events) if isinstance(e, _SseFinalMessage)), None)
usage_data: Final = next((e.usage for e in reversed(events) if isinstance(e, _SseUsage)), None)
content: Final = (
self._extract_content_from_message(final_message)
if final_message
else "".join(e.text for e in events if isinstance(e, _SseTextDelta))
)
return AgentCoreParsedResponse(content=content, usage=usage_data, final_message=final_message)
@ -533,103 +598,17 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
response: httpx.Response,
model: str,
):
"""
Internal sync generator that parses SSE and yields ModelResponse chunks.
"""
"""Internal sync generator that parses SSE and yields ModelResponse chunks."""
buffer = ""
for text_chunk in response.iter_text():
buffer += text_chunk
# Process complete lines
while "\n" in buffer:
line, buffer = buffer.split("\n", 1)
line = line.strip()
yield from self._line_to_chunks(line, model)
if not line or not line.startswith("data:"):
continue
json_str = line[5:].strip()
if not json_str:
continue
try:
data_obj = json.loads(json_str)
if not isinstance(data_obj, dict):
continue
# Process contentBlockDelta events
if "event" in data_obj and isinstance(data_obj["event"], dict):
event_payload = data_obj["event"]
content_block_delta = event_payload.get("contentBlockDelta")
if content_block_delta:
delta = content_block_delta.get("delta", {})
text = delta.get("text", "")
if text:
chunk = ModelResponseStream(
id=f"chatcmpl-{uuid.uuid4()}",
created=0,
model=model,
object="chat.completion.chunk",
)
chunk.choices = [
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content=text, role="assistant"),
)
]
yield chunk
# Process metadata/usage
metadata = event_payload.get("metadata")
if metadata and "usage" in metadata:
chunk = ModelResponseStream(
id=f"chatcmpl-{uuid.uuid4()}",
created=0,
model=model,
object="chat.completion.chunk",
)
chunk.choices = [
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(),
)
]
usage_data: AgentCoreUsage = metadata["usage"]
setattr(
chunk,
"usage",
Usage(
prompt_tokens=usage_data.get("inputTokens", 0),
completion_tokens=usage_data.get("outputTokens", 0),
total_tokens=usage_data.get("totalTokens", 0),
),
)
yield chunk
# Process final message
if "message" in data_obj and isinstance(data_obj["message"], dict):
chunk = ModelResponseStream(
id=f"chatcmpl-{uuid.uuid4()}",
created=0,
model=model,
object="chat.completion.chunk",
)
chunk.choices = [
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(),
)
]
yield chunk
except json.JSONDecodeError:
verbose_logger.debug("Skipping non-JSON SSE line: %s", line[:100])
continue
if buffer:
yield from self._line_to_chunks(buffer, model)
def get_sync_custom_stream_wrapper(
self,
@ -752,103 +731,19 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
response: httpx.Response,
model: str,
) -> AsyncGenerator[ModelResponseStream, None]:
"""
Internal async generator that parses SSE and yields ModelResponse chunks.
"""
"""Internal async generator that parses SSE and yields ModelResponse chunks."""
buffer = ""
async for text_chunk in response.aiter_text():
buffer += text_chunk
# Process complete lines
while "\n" in buffer:
line, buffer = buffer.split("\n", 1)
line = line.strip()
for chunk in self._line_to_chunks(line, model):
yield chunk
if not line or not line.startswith("data:"):
continue
json_str = line[5:].strip()
if not json_str:
continue
try:
data_obj = json.loads(json_str)
if not isinstance(data_obj, dict):
continue
# Process contentBlockDelta events
if "event" in data_obj and isinstance(data_obj["event"], dict):
event_payload = data_obj["event"]
content_block_delta = event_payload.get("contentBlockDelta")
if content_block_delta:
delta = content_block_delta.get("delta", {})
text = delta.get("text", "")
if text:
chunk = ModelResponseStream(
id=f"chatcmpl-{uuid.uuid4()}",
created=0,
model=model,
object="chat.completion.chunk",
)
chunk.choices = [
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content=text, role="assistant"),
)
]
yield chunk
# Process metadata/usage
metadata = event_payload.get("metadata")
if metadata and "usage" in metadata:
chunk = ModelResponseStream(
id=f"chatcmpl-{uuid.uuid4()}",
created=0,
model=model,
object="chat.completion.chunk",
)
chunk.choices = [
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(),
)
]
usage_data: AgentCoreUsage = metadata["usage"]
setattr(
chunk,
"usage",
Usage(
prompt_tokens=usage_data.get("inputTokens", 0),
completion_tokens=usage_data.get("outputTokens", 0),
total_tokens=usage_data.get("totalTokens", 0),
),
)
yield chunk
# Process final message
if "message" in data_obj and isinstance(data_obj["message"], dict):
chunk = ModelResponseStream(
id=f"chatcmpl-{uuid.uuid4()}",
created=0,
model=model,
object="chat.completion.chunk",
)
chunk.choices = [
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(),
)
]
yield chunk
except json.JSONDecodeError:
verbose_logger.debug("Skipping non-JSON SSE line: %s", line[:100])
continue
if buffer:
for chunk in self._line_to_chunks(buffer, model):
yield chunk
async def get_async_custom_stream_wrapper(
self,

View file

@ -20,6 +20,28 @@ import litellm
from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig
_AGENTCORE_MODEL = "bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent"
def _make_sync_sse_response(body: str) -> Mock:
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"}
mock_response.iter_text = Mock(return_value=iter([body]))
return mock_response
def _make_async_sse_response(body: str) -> Mock:
async def _aiter_text():
yield body
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"}
mock_response.aiter_text = Mock(return_value=_aiter_text())
return mock_response
class TestAgentCoreAcceptHeader:
"""Tests for Accept header in AgentCore requests."""
@ -641,3 +663,203 @@ class TestAgentCoreMultimodalContent:
payload = config.transform_request(messages=messages, **kwargs)
assert payload["content"] == content
assert payload["content"] is not content
class TestAgentCoreStringSsePayloads:
"""Regression tests for issue #25691: string SSE payloads from Strands agents."""
@pytest.fixture
def config(self):
return AmazonAgentCoreConfig()
def _complete_sync(self, body: str) -> str:
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(client, "post", return_value=_make_sync_sse_response(body)):
response = litellm.completion(
model=_AGENTCORE_MODEL,
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
api_key="test-jwt-token",
)
return "".join(
c.choices[0].delta.content
for c in response
if c.choices[0].delta.content
)
def _complete_non_streaming(self, body: str):
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"}
mock_response.text = body
with patch.object(client, "post", return_value=mock_response):
return litellm.completion(
model=_AGENTCORE_MODEL,
messages=[{"role": "user", "content": "hi"}],
client=client,
api_key="test-jwt-token",
)
async def _complete_async(self, body: str) -> str:
from unittest.mock import AsyncMock
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
client = AsyncHTTPHandler()
with patch.object(
client,
"post",
new_callable=AsyncMock,
return_value=_make_async_sse_response(body),
):
response = await litellm.acompletion(
model=_AGENTCORE_MODEL,
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
api_key="test-jwt-token",
)
content = ""
async for c in response:
if c.choices[0].delta.content:
content += c.choices[0].delta.content
return content
def test_sync_stream_quoted_string_payloads(self):
assert self._complete_sync('data: "hello"\ndata: " world"\n') == "hello world"
async def test_async_stream_quoted_string_payloads(self):
content = await self._complete_async('data: "hello"\ndata: " world"\n')
assert content == "hello world"
def test_sync_stream_plain_non_json_payload(self):
assert self._complete_sync("data: hello\n") == "hello"
async def test_async_stream_plain_non_json_payload(self):
assert await self._complete_async("data: hello\n") == "hello"
def test_sync_stream_skips_bare_done_sentinel(self):
assert self._complete_sync("data: hello\ndata: [DONE]\n") == "hello"
def test_sync_stream_skips_quoted_done_sentinel(self):
assert self._complete_sync('data: "[DONE]"\n') == ""
def test_non_streaming_sse_aggregator_quoted_string(self, config):
parsed = config._parse_sse_stream('data: "hi there"')
assert parsed["content"] == "hi there"
def test_non_streaming_public_path_surfaces_content_and_tokens(self):
response = self._complete_non_streaming('data: "hello from agent"\n')
assert response.choices[0].message.content == "hello from agent"
assert response.usage.completion_tokens > 0
def test_sync_stream_skips_malformed_json_frame(self):
assert self._complete_sync('data: {"event":\n') == ""
def test_sync_stream_non_dict_event_is_skipped_not_crashed(self):
body = "data: " + json.dumps({"event": "not-a-dict"}) + "\n"
assert self._complete_sync(body) == ""
async def test_async_stream_non_dict_event_is_skipped_not_crashed(self):
body = "data: " + json.dumps({"event": "not-a-dict"}) + "\n"
assert await self._complete_async(body) == ""
def test_sync_stream_flushes_unterminated_final_line(self):
assert self._complete_sync('data: "tail"') == "tail"
async def test_async_stream_flushes_unterminated_final_line(self):
assert await self._complete_async('data: "tail"') == "tail"
@staticmethod
def _content_block_delta_body(text: str) -> str:
return (
"data: "
+ json.dumps({"event": {"contentBlockDelta": {"delta": {"text": text}}}})
+ "\n"
)
@staticmethod
def _usage_event_body() -> str:
return (
"data: "
+ json.dumps(
{
"event": {
"metadata": {
"usage": {
"inputTokens": 5,
"outputTokens": 7,
"totalTokens": 12,
}
}
}
}
)
+ "\n"
)
@staticmethod
def _usage_chunk_matches(chunks) -> bool:
return any(
(u := getattr(c, "usage", None)) is not None
and u.prompt_tokens == 5
and u.completion_tokens == 7
and u.total_tokens == 12
for c in chunks
)
def test_sync_stream_dict_content_block_delta_still_works(self):
assert (
self._complete_sync(self._content_block_delta_body("dict path"))
== "dict path"
)
async def test_async_stream_dict_content_block_delta_still_works(self):
content = await self._complete_async(
self._content_block_delta_body("dict path")
)
assert content == "dict path"
def test_sync_generator_usage_event_still_surfaces_usage(self, config):
chunks = list(
config._stream_agentcore_response_sync(
_make_sync_sse_response(self._usage_event_body()), "test-model"
)
)
assert self._usage_chunk_matches(chunks)
async def test_async_generator_usage_event_still_surfaces_usage(self, config):
chunks = [
c
async for c in config._stream_agentcore_response(
_make_async_sse_response(self._usage_event_body()), "test-model"
)
]
assert self._usage_chunk_matches(chunks)
def test_dict_content_block_delta_empty_text_yields_no_content(self):
assert self._complete_sync(self._content_block_delta_body("")) == ""
def test_non_streaming_final_message_wins_over_deltas(self, config):
body = (
self._content_block_delta_body("delta text")
+ "data: "
+ json.dumps(
{"message": {"role": "assistant", "content": [{"text": "final text"}]}}
)
+ "\n"
)
parsed = config._parse_sse_stream(body)
assert parsed["content"] == "final text"
def test_sync_stream_non_string_non_dict_yields_no_content(self):
assert self._complete_sync("data: 123\ndata: [1, 2]\n") == ""
def test_sync_stream_empty_string_payload_yields_no_content(self):
assert self._complete_sync('data: ""\n') == ""