mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-22 00:31:44 +00:00
Merge 5968c18f73 into 252c71c0b2
This commit is contained in:
commit
e0d66ba188
2 changed files with 353 additions and 236 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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') == ""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue