This commit is contained in:
mubashir1osmani 2026-08-27 16:37:09 -04:00 • committed by GitHub
commit 68dfad4c73
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 1760 additions and 183 deletions

View file

@ -65,6 +65,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3
# Copy full source tree
@ -86,6 +87,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -63,6 +63,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3
# Copy full source tree
@ -84,6 +85,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -23,6 +23,7 @@ from .litellm_logging import Logging as LiteLLMLogging
if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
CLIENT_CONNECTION_CLASS = ClientConnection
@ -99,18 +100,18 @@ class RealTimeStreaming:
def __init__(
self,
websocket: Any,
backend_ws: CLIENT_CONNECTION_CLASS,
backend_ws: CLIENT_CONNECTION_CLASS | None,
logging_obj: LiteLLMLogging,
provider_config: BaseRealtimeConfig | None = None,
model: str = "",
user_api_key_dict: object | None = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
backend_uses_beta_protocol: bool | None = None,
force_transcription_model: str | None = None,
event_normalizer: RealtimeEventNormalizer | None = None,
):
self.websocket: _ClientWebSocket = websocket
self.backend_ws = backend_ws
self._backend_ws = backend_ws
self.logging_obj = logging_obj
self.messages: list[OpenAIRealtimeEvents] = []
self.input_message: dict = {}
@ -202,6 +203,19 @@ class RealTimeStreaming:
"output_audio": "audio",
}
@property
def backend_ws(self) -> CLIENT_CONNECTION_CLASS:
"""
The backend websocket, for the forwarding paths that require one.
Providers that stream over a non-websocket transport (Bedrock uses the AWS SDK
bidirectional stream) construct this class only for its message store and spend
logging, and pass ``backend_ws=None``; reaching a forwarding path from there is a bug.
"""
if self._backend_ws is None:
raise RuntimeError("RealTimeStreaming was constructed without a backend websocket")
return self._backend_ws
def _should_store_message(
self,
message_obj: dict | OpenAIRealtimeEvents,

View file

@ -2,23 +2,32 @@
This file contains the handler for AWS Bedrock Nova Sonic realtime API.
This uses aws_sdk_bedrock_runtime for bidirectional streaming with Nova Sonic.
Spend / budget logging follows the same RealTimeStreaming path as OpenAI/Azure:
store_message for backend events, store_input for client events, log_messages on close.
"""
import asyncio
import contextlib
import json
from typing import Any, Final
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from pydantic import TypeAdapter
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError
from .transformation import BedrockRealtimeConfig
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
_CLIENT_MODALITIES_ADAPTER: Final[TypeAdapter["list[str] | None"]] = TypeAdapter(list[str] | None)
_EMPTY_METADATA: Final[Mapping[str, object]] = MappingProxyType({})
class BedrockRealtime(BaseAWSLLM):
@ -46,7 +55,9 @@ class BedrockRealtime(BaseAWSLLM):
aws_sts_endpoint: str | None = None,
aws_bedrock_runtime_endpoint: str | None = None,
aws_external_id: str | None = None,
**kwargs,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
litellm_metadata: Mapping[str, object] | None = None,
**kwargs: object,
):
"""
Establish bidirectional streaming connection with Bedrock Nova Sonic.
@ -120,6 +131,33 @@ class BedrockRealtime(BaseAWSLLM):
transformation_config: Final = BedrockRealtimeConfig()
pre_call_args: Final = MappingProxyType(
{
"api_base": endpoint_uri,
"complete_input_dict": MappingProxyType({"model": model}),
}
)
logging_obj.pre_call(
input=None,
api_key=api_key or "",
additional_args=dict(pre_call_args), # mutable-ok: Logging.pre_call expects a mutable dict
)
# RealTimeStreaming owns spend logging for other realtime providers. Bedrock cannot
# use its WebSocket bidirectional_forward (AWS SDK stream instead), but store_message /
# store_input / log_messages are the same path used by OpenAI and Azure.
request_data: Final = MappingProxyType(
{"litellm_metadata": litellm_metadata if litellm_metadata is not None else _EMPTY_METADATA}
)
realtime_streaming: Final = RealTimeStreaming(
websocket=websocket,
backend_ws=None, # Bedrock streams over the AWS SDK; only store/log are used here
logging_obj=logging_obj,
model=model,
user_api_key_dict=user_api_key_dict,
request_data=dict(request_data), # mutable-ok: RealTimeStreaming stores request_data as dict
)
try:
# Initialize the bidirectional stream
bedrock_stream: Final = await bedrock_client.invoke_model_with_bidirectional_stream(
@ -128,7 +166,9 @@ class BedrockRealtime(BaseAWSLLM):
verbose_proxy_logger.debug("Bedrock Realtime: Bidirectional stream established")
await websocket.send_text(json.dumps(transformation_config.session_created_event(model, logging_obj)))
session_created: Final = transformation_config.session_created_event(model, logging_obj)
realtime_streaming.store_message(session_created)
await websocket.send_text(json.dumps(session_created))
verbose_proxy_logger.debug("Bedrock Realtime: sent session.created to client on connect")
# Track state for transformation
@ -142,35 +182,43 @@ class BedrockRealtime(BaseAWSLLM):
"session_configuration_request": None,
}
# Create tasks for bidirectional forwarding
client_to_bedrock_task: Final = asyncio.create_task(
self._forward_client_to_bedrock(
websocket,
bedrock_stream,
transformation_config,
model,
session_state,
logging_obj,
try:
client_to_bedrock_task: Final = asyncio.create_task(
self._forward_client_to_bedrock(
websocket,
bedrock_stream,
transformation_config,
model,
session_state,
logging_obj,
realtime_streaming,
)
)
)
bedrock_to_client_task: Final = asyncio.create_task(
self._forward_bedrock_to_client(
bedrock_stream,
websocket,
transformation_config,
model,
logging_obj,
session_state,
bedrock_to_client_task: Final = asyncio.create_task(
self._forward_bedrock_to_client(
bedrock_stream,
websocket,
transformation_config,
model,
logging_obj,
session_state,
realtime_streaming,
)
)
)
# Wait for both tasks to complete
await asyncio.gather(
client_to_bedrock_task,
bedrock_to_client_task,
return_exceptions=True,
)
await asyncio.gather(
client_to_bedrock_task,
bedrock_to_client_task,
return_exceptions=True,
)
finally:
for pending_usage_event in transformation_config.flush_pending_usage_as_response_done(
session_state.get("current_response_id"),
session_state.get("current_conversation_id"),
):
realtime_streaming.store_message(pending_usage_event)
await realtime_streaming.log_messages()
except Exception as e:
verbose_proxy_logger.exception("Error in BedrockRealtime.async_realtime: %s", e)
@ -188,6 +236,7 @@ class BedrockRealtime(BaseAWSLLM):
model: str,
session_state: dict,
logging_obj: LiteLLMLogging | None = None,
realtime_streaming: RealTimeStreaming | None = None,
):
"""Forward messages from client WebSocket to Bedrock stream."""
from aws_sdk_bedrock_runtime.models import (
@ -208,6 +257,9 @@ class BedrockRealtime(BaseAWSLLM):
message = await client_ws.receive_text()
verbose_proxy_logger.debug("Bedrock Realtime: Received from client: %s", message[:200])
if realtime_streaming is not None:
realtime_streaming.store_input(message)
# Transform OpenAI format to Bedrock format
transformed_messages = transformation_config.transform_realtime_request(
message=message,
@ -230,11 +282,12 @@ class BedrockRealtime(BaseAWSLLM):
parsed_client_message.get("session", {}).get("modalities")
)
if client_message_type == "session.update":
await client_ws.send_text(
json.dumps(
transformation_config.session_updated_event(model, logging_obj, requested_modalities)
)
session_updated = transformation_config.session_updated_event( # rebind-ok: per-iteration local
model, logging_obj, requested_modalities
)
if realtime_streaming is not None:
realtime_streaming.store_message(session_updated)
await client_ws.send_text(json.dumps(session_updated))
except Exception as e:
verbose_proxy_logger.debug("Client to Bedrock forwarding ended: %s", e, exc_info=True)
@ -252,6 +305,7 @@ class BedrockRealtime(BaseAWSLLM):
model: str,
logging_obj: LiteLLMLogging,
session_state: dict,
realtime_streaming: RealTimeStreaming | None = None,
):
"""Forward messages from Bedrock stream to client WebSocket."""
try:
@ -304,6 +358,8 @@ class BedrockRealtime(BaseAWSLLM):
# Send transformed messages to client
openai_messages = transformed_response.get("response", [])
for openai_message in openai_messages:
if realtime_streaming is not None:
realtime_streaming.store_message(openai_message)
message_json = json.dumps(openai_message)
await client_ws.send_text(message_json)
verbose_proxy_logger.debug("Bedrock Realtime: Sent to client: %s", message_json[:200])

View file

@ -7,7 +7,9 @@ Transforms between OpenAI Realtime API format and Bedrock Nova Sonic format.
import base64
import json
import uuid as uuid_lib
from typing import Any, Final
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel
@ -18,8 +20,11 @@ from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.bedrock.realtime.trigger_audio import ready_trigger_pcm
from litellm.types.llms.openai import (
OpenAIRealtimeContentPartDone,
OpenAIRealtimeConversationItemAdded,
OpenAIRealtimeDoneEvent,
OpenAIRealtimeEvents,
OpenAIRealtimeFunctionCallArgumentsDelta,
OpenAIRealtimeFunctionCallArgumentsDone,
OpenAIRealtimeOutputItemDone,
OpenAIRealtimeResponseAudioDone,
OpenAIRealtimeResponseContentPartAdded,
@ -27,6 +32,7 @@ from litellm.types.llms.openai import (
OpenAIRealtimeResponseDoneObject,
OpenAIRealtimeResponseTextDone,
OpenAIRealtimeStreamResponseBaseObject,
OpenAIRealtimeStreamResponseOutputItem,
OpenAIRealtimeStreamResponseOutputItemAdded,
OpenAIRealtimeStreamSession,
OpenAIRealtimeStreamSessionEvents,
@ -36,19 +42,127 @@ from litellm.types.realtime import (
RealtimeResponseTransformInput,
RealtimeResponseTypedDict,
)
from litellm.utils import get_empty_usage
class BedrockContentEnd(BaseModel):
stopReason: str | None = None
class BedrockToolUse(BaseModel):
toolUseId: str = ""
toolName: str = ""
content: object | None = None
input: object | None = None
def arguments(self) -> str:
"""Nova Sonic puts tool args in ``content`` as a JSON string; older payloads use ``input``."""
return json.dumps(_parse_bedrock_tool_use_input(self.content if self.content is not None else self.input))
TRIGGER_AUDIO_SAMPLE_RATE_HERTZ: Final = 16000
TRIGGER_AUDIO_BYTES_PER_SECOND: Final = TRIGGER_AUDIO_SAMPLE_RATE_HERTZ * 2
TRIGGER_LEADING_SILENCE: Final = bytes(TRIGGER_AUDIO_BYTES_PER_SECOND // 2)
TRIGGER_TRAILING_SILENCE: Final = bytes(TRIGGER_AUDIO_BYTES_PER_SECOND * 3)
TRIGGER_AUDIO_CHUNK_SIZE: Final = 1024
_EMPTY_USAGE_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
_USAGE_SNAPSHOT_KEYS: Final = (
"input_speech",
"input_text",
"output_speech",
"output_text",
"total_input",
"total_output",
"total",
)
def _parse_bedrock_tool_use_input(raw_input: object) -> object:
if not raw_input:
return {} # mutable-ok: tool args are JSON-serializable wire values
if not isinstance(raw_input, str):
return raw_input
try:
return json.loads(raw_input)
except json.JSONDecodeError:
return {} # mutable-ok: tool args are JSON-serializable wire values
def _as_nonneg_int(value: object) -> int:
if isinstance(value, bool) or not isinstance(value, (int, float)):
return 0
return max(0, int(value))
def _mapping_or_empty(value: object) -> Mapping[str, object]:
return value if isinstance(value, Mapping) else _EMPTY_USAGE_MAPPING
def _empty_usage_snapshot() -> Mapping[str, int]:
return MappingProxyType(
{
"input_speech": 0,
"input_text": 0,
"output_speech": 0,
"output_text": 0,
"total_input": 0,
"total_output": 0,
"total": 0,
}
)
def _usage_snapshot_from_event(usage_event: Mapping[str, object]) -> Mapping[str, int]:
details: Final = _mapping_or_empty(usage_event.get("details"))
total_block: Final = _mapping_or_empty(details.get("total"))
input_block: Final = _mapping_or_empty(total_block.get("input"))
output_block: Final = _mapping_or_empty(total_block.get("output"))
input_speech: Final = _as_nonneg_int(input_block.get("speechTokens"))
input_text: Final = _as_nonneg_int(input_block.get("textTokens"))
output_speech: Final = _as_nonneg_int(output_block.get("speechTokens"))
output_text: Final = _as_nonneg_int(output_block.get("textTokens"))
total_input: Final = _as_nonneg_int(usage_event.get("totalInputTokens")) or (input_speech + input_text)
total_output: Final = _as_nonneg_int(usage_event.get("totalOutputTokens")) or (output_speech + output_text)
total: Final = _as_nonneg_int(usage_event.get("totalTokens")) or (total_input + total_output)
return MappingProxyType(
{
"input_speech": input_speech,
"input_text": input_text,
"output_speech": output_speech,
"output_text": output_text,
"total_input": total_input,
"total_output": total_output,
"total": total,
}
)
def _usage_snapshot_delta(current: Mapping[str, int], previous: Mapping[str, int]) -> Mapping[str, int]:
return MappingProxyType({key: max(0, current.get(key, 0) - previous.get(key, 0)) for key in _USAGE_SNAPSHOT_KEYS})
def _openai_usage_from_snapshot(snapshot: Mapping[str, int]) -> dict[str, object]: # mutable-ok: OpenAI usage wire dict
input_tokens: Final = snapshot.get("total_input", 0) or (
snapshot.get("input_speech", 0) + snapshot.get("input_text", 0)
)
output_tokens: Final = snapshot.get("total_output", 0) or (
snapshot.get("output_speech", 0) + snapshot.get("output_text", 0)
)
return { # mutable-ok: OpenAI response.done usage is a JSON-serializable dict
"total_tokens": snapshot.get("total", 0) or (input_tokens + output_tokens),
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"input_token_details": { # mutable-ok: nested OpenAI usage wire shape
"text_tokens": snapshot.get("input_text", 0),
"audio_tokens": snapshot.get("input_speech", 0),
"cached_tokens": 0,
},
"output_token_details": { # mutable-ok: nested OpenAI usage wire shape
"text_tokens": snapshot.get("output_text", 0),
"audio_tokens": snapshot.get("output_speech", 0),
},
}
class BedrockRealtimeConfig(BaseRealtimeConfig):
"""Configuration for Bedrock Nova Sonic realtime transformations."""
@ -60,6 +174,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
self.audio_content_name = str(uuid_lib.uuid4())
self.prompt_started = False
self.client_audio_streamed = False
self._usage_totals = _empty_usage_snapshot()
self._usage_at_last_response_done = _empty_usage_snapshot()
# Default configuration values
# Inference configuration
@ -87,6 +203,32 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
# Text configuration
self.text_media_type = "text/plain"
def record_usage_event(self, usage_event: Mapping[str, object]) -> None:
self._usage_totals = _usage_snapshot_from_event(usage_event)
def has_unbilled_usage(self) -> bool:
delta: Final = _usage_snapshot_delta(self._usage_totals, self._usage_at_last_response_done)
return any(value > 0 for value in delta.values())
def consume_usage_for_response_done(self) -> dict[str, object]: # mutable-ok: OpenAI usage wire dict
delta: Final = _usage_snapshot_delta(self._usage_totals, self._usage_at_last_response_done)
self._usage_at_last_response_done = self._usage_totals
return _openai_usage_from_snapshot(delta)
def flush_pending_usage_as_response_done(
self,
current_response_id: str | None = None,
current_conversation_id: str | None = None,
) -> list[OpenAIRealtimeEvents]: # mutable-ok: callers store into mutable message lists
if not self.has_unbilled_usage():
return [] # mutable-ok: empty OpenAI event list for callers that append/extend
events, _, _, _ = self._response_done_events(
current_response_id,
current_conversation_id,
mint_ids_if_missing=True,
)
return list(events) # mutable-ok: session-close flush is stored into mutable message lists
def validate_environment(self, headers: dict, model: str, api_key: str | None = None) -> dict:
"""Validate environment - no special validation needed for Bedrock."""
return headers
@ -668,6 +810,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
current_response_id: str | None,
current_output_item_id: str | None,
current_conversation_id: str | None,
current_delta_type: ALL_DELTA_TYPES | None = None,
) -> tuple[
list[OpenAIRealtimeEvents],
str | None,
@ -678,14 +821,9 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
"""
Transform Bedrock contentStart event to OpenAI response events.
Args:
event: Bedrock contentStart event
current_response_id: Current response ID
current_output_item_id: Current output item ID
current_conversation_id: Current conversation ID
Returns:
Tuple of (events, response_id, output_item_id, conversation_id, delta_type)
Bedrock streams one content block at a time (TEXT, AUDIO, TOOL, …). Only
ASSISTANT blocks open an OpenAI response/item lifecycle. Non-assistant
blocks must not clobber in-flight assistant part state.
"""
content_start: Final = event["contentStart"]
role: Final = content_start.get("role")
@ -696,40 +834,37 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
current_response_id,
current_output_item_id,
current_conversation_id,
None,
current_delta_type,
)
verbose_logger.debug("Handling ASSISTANT contentStart")
# Initialize IDs if needed
is_new_response: Final = not current_response_id
if not current_response_id:
current_response_id = f"resp_{uuid.uuid4()}"
if not current_output_item_id:
current_output_item_id = f"item_{uuid.uuid4()}"
current_output_item_id = f"item_{uuid.uuid4()}"
if not current_conversation_id:
current_conversation_id = f"conv_{uuid.uuid4()}"
# Determine content type
content_type: Final = content_start.get("type", "TEXT")
current_delta_type: Final[ALL_DELTA_TYPES] = "text" if content_type == "TEXT" else "audio"
next_delta_type: Final[ALL_DELTA_TYPES] = "text" if content_type == "TEXT" else "audio"
returned_messages: Final[list[OpenAIRealtimeEvents]] = []
# Send response.created
response_created: Final = OpenAIRealtimeStreamResponseBaseObject(
type="response.created",
event_id=f"event_{uuid.uuid4()}",
response={
"object": "realtime.response",
"id": current_response_id,
"status": "in_progress",
"output": [],
"conversation_id": current_conversation_id,
},
)
returned_messages.append(response_created)
if is_new_response:
response_created: Final = OpenAIRealtimeStreamResponseBaseObject(
type="response.created",
event_id=f"event_{uuid.uuid4()}",
response={
"object": "realtime.response",
"id": current_response_id,
"status": "in_progress",
"output": [],
"conversation_id": current_conversation_id,
},
)
returned_messages.append(response_created)
# Send response.output_item.added
output_item_added: Final = OpenAIRealtimeStreamResponseOutputItemAdded(
type="response.output_item.added",
response_id=current_response_id,
@ -745,16 +880,13 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
)
returned_messages.append(output_item_added)
# Send response.content_part.added
content_part_added: Final = OpenAIRealtimeResponseContentPartAdded(
type="response.content_part.added",
content_index=0,
output_index=0,
event_id=f"event_{uuid.uuid4()}",
item_id=current_output_item_id,
part=(
{"type": "text", "text": ""} if current_delta_type == "text" else {"type": "audio", "transcript": ""}
),
part=({"type": "text", "text": ""} if next_delta_type == "text" else {"type": "audio", "transcript": ""}),
response_id=current_response_id,
)
returned_messages.append(content_part_added)
@ -764,7 +896,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
current_response_id,
current_output_item_id,
current_conversation_id,
current_delta_type,
next_delta_type,
)
def transform_text_output_event(
@ -869,7 +1001,10 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
verbose_logger.debug("Handling contentEnd: %s", content_end)
if not current_output_item_id or not current_response_id:
return [], current_delta_chunks
return [], None
if content_end.get("type") == "TOOL" or current_delta_type not in ("text", "audio"):
return [], None
returned_messages: Final[list[OpenAIRealtimeEvents]] = []
@ -970,40 +1105,42 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
Tuple of (events, reset_output_item_id, reset_response_id, reset_delta_type)
"""
verbose_logger.debug("Handling promptEnd")
return self._response_done_events(current_response_id, current_conversation_id)
return self._response_done_events(
current_response_id,
current_conversation_id,
mint_ids_if_missing=self.has_unbilled_usage(),
)
def _response_done_events(
self,
current_response_id: str | None,
current_conversation_id: str | None,
*,
mint_ids_if_missing: bool = False,
) -> tuple[
list[OpenAIRealtimeEvents],
str | None,
str | None,
ALL_DELTA_TYPES | None,
]:
if not current_response_id or not current_conversation_id:
response_id: Final = current_response_id or (f"resp_{uuid.uuid4()}" if mint_ids_if_missing else None)
conversation_id: Final = current_conversation_id or (f"conv_{uuid.uuid4()}" if mint_ids_if_missing else None)
if not response_id or not conversation_id:
return [], None, None, None
usage_obj: Final = get_empty_usage()
response_done: Final = OpenAIRealtimeDoneEvent(
type="response.done",
event_id=f"event_{uuid.uuid4()}",
response=OpenAIRealtimeResponseDoneObject(
object="realtime.response",
id=current_response_id,
id=response_id,
status="completed",
output=[],
conversation_id=current_conversation_id,
usage={
"prompt_tokens": usage_obj.prompt_tokens,
"completion_tokens": usage_obj.completion_tokens,
"total_tokens": usage_obj.total_tokens,
},
conversation_id=conversation_id,
usage=self.consume_usage_for_response_done(),
),
)
# Reset state for next response
return [response_done], None, None, None
def transform_tool_use_event(
@ -1011,55 +1148,126 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
event: dict,
current_output_item_id: str | None,
current_response_id: str | None,
conversation_id: str,
) -> tuple[list[OpenAIRealtimeEvents], str, str]:
"""
Transform Bedrock toolUse event to OpenAI format.
Transform a Bedrock toolUse event into the full OpenAI function-call lifecycle.
Args:
event: Bedrock toolUse event
current_output_item_id: Current output item ID
current_response_id: Current response ID
Nova Sonic delivers one tool call, fully formed, in a single event, and starts the
block with ``contentStart`` role ``TOOL``, which opens no OpenAI response. Mint the
response/item ids when they are missing, opening the response first so its
``response.done`` is never unmatched, emit the item added/delta/done trio the OpenAI
realtime protocol requires around ``function_call_arguments.done``, then close the
response so nothing is left in progress and downstream spend logging can harvest the
call from ``response.done`` output.
Returns:
Tuple of (events, tool_call_id, tool_name) for tracking
Tuple of (events, tool_call_id, tool_name). The caller clears the response and
item ids, since this sequence closes the response it emits.
"""
verbose_logger.debug("Handling toolUse")
tool_use: Final = event["toolUse"]
tool_use: Final = BedrockToolUse.model_validate(event["toolUse"])
if not current_output_item_id or not current_response_id:
return [], "", ""
is_new_response: Final = not current_response_id
response_id: Final = current_response_id or f"resp_{uuid.uuid4()}"
item_id: Final = current_output_item_id or f"item_{uuid.uuid4()}"
tool_call_id: Final = tool_use.toolUseId
tool_name: Final = tool_use.toolName
arguments: Final = tool_use.arguments()
# Parse the tool input
tool_input = {}
if "input" in tool_use:
try:
tool_input = json.loads(tool_use["input"]) if isinstance(tool_use["input"], str) else tool_use["input"]
except json.JSONDecodeError:
tool_input = {}
tool_call_id: Final = tool_use.get("toolUseId", "")
tool_name: Final = tool_use.get("toolName", "")
# Create a function call arguments done event
# This is a custom event format that matches what clients expect
from typing import cast
function_call_event: Final[dict[str, Any]] = {
"type": "response.function_call_arguments.done",
"event_id": f"event_{uuid.uuid4()}",
"response_id": current_response_id,
"item_id": current_output_item_id,
"output_index": 0,
"call_id": tool_call_id,
"name": tool_name,
"arguments": json.dumps(tool_input),
}
return (
[cast(OpenAIRealtimeEvents, function_call_event)],
tool_call_id,
tool_name,
function_call_item: Final = OpenAIRealtimeStreamResponseOutputItem(
id=item_id,
object="realtime.item",
type="function_call",
status="completed",
call_id=tool_call_id,
name=tool_name,
arguments=arguments,
)
pending_item: Final = OpenAIRealtimeStreamResponseOutputItem(
{**function_call_item, "status": "in_progress", "arguments": ""}
)
# A tool turn that Bedrock opens with contentStart role TOOL has no response yet, so the
# response.done below would close an id the client never saw opened.
response_created: Final[tuple[OpenAIRealtimeEvents, ...]] = (
(
OpenAIRealtimeStreamResponseBaseObject(
type="response.created",
event_id=f"event_{uuid.uuid4()}",
response={
"object": "realtime.response",
"id": response_id,
"status": "in_progress",
"output": [],
"conversation_id": conversation_id,
},
),
)
if is_new_response
else ()
)
events: Final[list[OpenAIRealtimeEvents]] = [
*response_created,
OpenAIRealtimeStreamResponseOutputItemAdded(
type="response.output_item.added",
event_id=f"event_{uuid.uuid4()}",
response_id=response_id,
output_index=0,
item=pending_item,
),
# Pipecat registers call_id from conversation.item.added; without it the
# function_call_arguments.done below is dropped as an unknown call.
OpenAIRealtimeConversationItemAdded(
type="conversation.item.added",
event_id=f"event_{uuid.uuid4()}",
previous_item_id=None,
item=pending_item,
),
# Nova Sonic delivers the whole argument payload at once; emit one delta anyway
# so clients that accumulate deltas rather than read `.done` still get the args.
OpenAIRealtimeFunctionCallArgumentsDelta(
type="response.function_call_arguments.delta",
event_id=f"event_{uuid.uuid4()}",
response_id=response_id,
item_id=item_id,
output_index=0,
call_id=tool_call_id,
delta=arguments,
),
OpenAIRealtimeFunctionCallArgumentsDone(
type="response.function_call_arguments.done",
event_id=f"event_{uuid.uuid4()}",
response_id=response_id,
item_id=item_id,
output_index=0,
call_id=tool_call_id,
name=tool_name,
arguments=arguments,
),
OpenAIRealtimeOutputItemDone(
type="response.output_item.done",
event_id=f"event_{uuid.uuid4()}",
response_id=response_id,
output_index=0,
item=function_call_item,
),
OpenAIRealtimeDoneEvent(
type="response.done",
event_id=f"event_{uuid.uuid4()}",
response=OpenAIRealtimeResponseDoneObject(
object="realtime.response",
id=response_id,
status="completed",
output=[function_call_item],
conversation_id=conversation_id,
usage=self.consume_usage_for_response_done(),
),
),
]
return events, tool_call_id, tool_name
def transform_conversation_item_create_tool_result_event(self, json_message: dict) -> list[str]:
"""
@ -1161,13 +1369,25 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
"session_configuration_request": realtime_response_transform_input.get("session_configuration_request"),
}
# Extract state
current_output_item_id = realtime_response_transform_input.get("current_output_item_id")
current_response_id = realtime_response_transform_input.get("current_response_id")
current_conversation_id = realtime_response_transform_input.get("current_conversation_id")
current_delta_chunks = realtime_response_transform_input.get("current_delta_chunks")
current_delta_type = realtime_response_transform_input.get("current_delta_type")
session_configuration_request = realtime_response_transform_input.get("session_configuration_request")
# Extract state. Session state is intentionally re-bound as each Bedrock event is folded in.
current_output_item_id = realtime_response_transform_input.get(
"current_output_item_id"
) # rebind-ok: session state machine
current_response_id = realtime_response_transform_input.get(
"current_response_id"
) # rebind-ok: session state machine
current_conversation_id = realtime_response_transform_input.get(
"current_conversation_id"
) # rebind-ok: session state machine
current_delta_chunks = realtime_response_transform_input.get(
"current_delta_chunks"
) # rebind-ok: session state machine
current_delta_type = realtime_response_transform_input.get(
"current_delta_type"
) # rebind-ok: session state machine
session_configuration_request = realtime_response_transform_input.get(
"session_configuration_request"
) # rebind-ok: session state machine
returned_messages: Final[list[OpenAIRealtimeEvents]] = []
@ -1176,25 +1396,46 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
# Route to appropriate transformation method
if "sessionStart" in event:
session_configuration_request = json.dumps({"configured": True})
session_configuration_request = json.dumps({"configured": True}) # rebind-ok: session state machine
elif "usageEvent" in event:
usage_event: Final = event["usageEvent"]
if isinstance(usage_event, dict):
self.record_usage_event(usage_event)
if current_response_id is None and self.has_unbilled_usage():
(
done_events,
current_output_item_id, # rebind-ok: session state machine
current_response_id, # rebind-ok: session state machine
current_delta_type, # rebind-ok: session state machine
) = self._response_done_events(
None,
current_conversation_id,
mint_ids_if_missing=True,
)
returned_messages.extend(done_events)
current_delta_chunks = None # rebind-ok: session state machine
elif "contentStart" in event:
(
events,
current_response_id,
current_output_item_id,
current_conversation_id,
current_delta_type,
current_response_id, # rebind-ok: session state machine
current_output_item_id, # rebind-ok: session state machine
current_conversation_id, # rebind-ok: session state machine
current_delta_type, # rebind-ok: session state machine
) = self.transform_content_start_event(
event,
current_response_id,
current_output_item_id,
current_conversation_id,
current_delta_type,
)
returned_messages.extend(events)
if events:
current_delta_chunks = None # rebind-ok: session state machine
elif "textOutput" in event:
events, current_delta_chunks = self.transform_text_output_event(
events, current_delta_chunks = self.transform_text_output_event( # rebind-ok: session state machine
event,
current_output_item_id,
current_response_id,
@ -1203,11 +1444,14 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
returned_messages.extend(events)
elif "audioOutput" in event:
events = self.transform_audio_output_event(event, current_output_item_id, current_response_id)
events = self.transform_audio_output_event(
event, current_output_item_id, current_response_id
) # rebind-ok: session state machine
returned_messages.extend(events)
elif "contentEnd" in event:
events, current_delta_chunks = self.transform_content_end_event(
content_end: Final = event["contentEnd"]
events, current_delta_chunks = self.transform_content_end_event( # rebind-ok: session state machine
event,
current_output_item_id,
current_response_id,
@ -1215,31 +1459,48 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
current_delta_chunks,
)
returned_messages.extend(events)
if BedrockContentEnd.model_validate(event["contentEnd"]).stopReason == "END_TURN":
current_delta_chunks = None # rebind-ok: session state machine
current_delta_type = None # rebind-ok: session state machine
current_output_item_id = None # rebind-ok: session state machine
is_end_turn: Final = BedrockContentEnd.model_validate(content_end).stopReason == "END_TURN"
if is_end_turn:
(
done_events,
current_output_item_id,
current_response_id,
current_delta_type,
current_output_item_id, # rebind-ok: session state machine
current_response_id, # rebind-ok: session state machine
current_delta_type, # rebind-ok: session state machine
) = self._response_done_events(current_response_id, current_conversation_id)
returned_messages.extend(done_events)
current_delta_chunks = None # rebind-ok: session state machine
elif "toolUse" in event:
current_conversation_id = ( # rebind-ok: session state machine
current_conversation_id or f"conv_{uuid.uuid4()}"
)
events, tool_call_id, tool_name = self.transform_tool_use_event(
event, current_output_item_id, current_response_id
event,
current_output_item_id,
current_response_id,
current_conversation_id,
)
returned_messages.extend(events)
# Store tool call info for potential use
# transform_tool_use_event closes the response it emits, so the tool ids must not
# survive into the post-tool assistant turn.
current_output_item_id = None # rebind-ok: session state machine
current_response_id = None # rebind-ok: session state machine
current_delta_chunks = None # rebind-ok: session state machine
current_delta_type = None # rebind-ok: session state machine
verbose_logger.debug("Tool use event: %s (ID: %s)", tool_name, tool_call_id)
elif "promptEnd" in event or "completionEnd" in event:
(
events,
current_output_item_id,
current_response_id,
current_delta_type,
current_output_item_id, # rebind-ok: session state machine
current_response_id, # rebind-ok: session state machine
current_delta_type, # rebind-ok: session state machine
) = self.transform_prompt_end_event(event, current_response_id, current_conversation_id)
returned_messages.extend(events)
current_delta_chunks = None # rebind-ok: session state machine
return {
"response": returned_messages,

View file

@ -51318,6 +51318,62 @@
}
]
},
"amazon.nova-sonic-v1:0": {
"input_cost_per_audio_token": 3.4e-06,
"input_cost_per_token": 6e-08,
"litellm_provider": "bedrock",
"max_input_tokens": 300000,
"max_output_tokens": 10000,
"max_tokens": 10000,
"mode": "realtime",
"output_cost_per_audio_token": 1.36e-05,
"output_cost_per_token": 2.4e-07,
"source": "https://aws.amazon.com/bedrock/pricing/",
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"amazon.nova-2-sonic-v1:0": {
"input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 3.3e-07,
"litellm_provider": "bedrock",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "realtime",
"output_cost_per_audio_token": 1.2e-05,
"output_cost_per_token": 2.75e-06,
"source": "https://aws.amazon.com/bedrock/pricing/",
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gemini/gemini-3.5-live-translate-preview": {
"input_cost_per_audio_token": 3.5e-06,
"input_cost_per_token": 3.5e-06,

View file

@ -482,6 +482,8 @@ async def _arealtime(
aws_sts_endpoint=aws_sts_endpoint,
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
aws_external_id=aws_external_id,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
)
elif _custom_llm_provider == "xai":
api_base = (

View file

@ -2107,6 +2107,16 @@ class OpenAIRealtimeContentPartDone(TypedDict):
type: Literal["response.content_part.done"]
class OpenAIRealtimeFunctionCallArgumentsDelta(TypedDict):
type: ReadOnly[Literal["response.function_call_arguments.delta"]]
event_id: ReadOnly[str]
response_id: ReadOnly[str]
item_id: ReadOnly[str]
output_index: ReadOnly[int]
call_id: ReadOnly[str]
delta: ReadOnly[str]
class OpenAIRealtimeFunctionCallArgumentsDone(TypedDict):
type: Literal["response.function_call_arguments.done"]
event_id: str
@ -2183,6 +2193,7 @@ OpenAIRealtimeEvents = (
| OpenAIRealtimeResponseAudioDone
| OpenAIRealtimeContentPartDone
| OpenAIRealtimeOutputItemDone
| OpenAIRealtimeFunctionCallArgumentsDelta
| OpenAIRealtimeFunctionCallArgumentsDone
| OpenAIRealtimeDoneEvent
)

View file

@ -51318,6 +51318,62 @@
}
]
},
"amazon.nova-sonic-v1:0": {
"input_cost_per_audio_token": 3.4e-06,
"input_cost_per_token": 6e-08,
"litellm_provider": "bedrock",
"max_input_tokens": 300000,
"max_output_tokens": 10000,
"max_tokens": 10000,
"mode": "realtime",
"output_cost_per_audio_token": 1.36e-05,
"output_cost_per_token": 2.4e-07,
"source": "https://aws.amazon.com/bedrock/pricing/",
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"amazon.nova-2-sonic-v1:0": {
"input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 3.3e-07,
"litellm_provider": "bedrock",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "realtime",
"output_cost_per_audio_token": 1.2e-05,
"output_cost_per_token": 2.75e-06,
"source": "https://aws.amazon.com/bedrock/pricing/",
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gemini/gemini-3.5-live-translate-preview": {
"input_cost_per_audio_token": 3.5e-06,
"input_cost_per_token": 3.5e-06,

View file

@ -30,6 +30,21 @@ def _make_transcript_event(text: str, item_id: str = "item_x") -> bytes:
).encode()
def test_store_and_log_work_without_a_backend_websocket():
"""
Providers that stream over a non-websocket transport (Bedrock's AWS SDK stream) build this
class only to store and log; store/log must work with backend_ws=None, and any forwarding
path reached from there must fail loudly rather than on a placeholder object.
"""
streaming = RealTimeStreaming(MagicMock(), None, MagicMock())
streaming.store_message(json.dumps({"type": "session.created", "session": {"id": "s"}}))
assert [message["type"] for message in streaming.messages] == ["session.created"]
with pytest.raises(RuntimeError, match="without a backend websocket"):
_ = streaming.backend_ws
def test_realtime_streaming_store_message():
# Setup
websocket = MagicMock()
@ -1403,7 +1418,6 @@ async def test_realtime_guardrail_blocks_prompt_injection(monkeypatch: pytest.Mo
)
@pytest.mark.asyncio
async def test_realtime_guardrail_allows_clean_transcript(monkeypatch: pytest.MonkeyPatch):
"""
@ -1460,7 +1474,6 @@ async def test_realtime_guardrail_allows_clean_transcript(monkeypatch: pytest.Mo
assert len(response_creates) == 1, f"Clean transcript should trigger response.create, got: {sent_to_backend}"
@pytest.mark.asyncio
async def test_realtime_text_input_guardrail_blocks_and_returns_error(monkeypatch: pytest.MonkeyPatch):
"""
@ -1554,7 +1567,6 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error(monkeypatc
assert len(original_items) == 0, f"Blocked item should not be forwarded to backend, got: {original_items}"
@pytest.mark.asyncio
async def test_realtime_function_call_output_guardrail_blocks_and_returns_error(monkeypatch: pytest.MonkeyPatch):
"""
@ -1643,7 +1655,6 @@ async def test_realtime_function_call_output_guardrail_blocks_and_returns_error(
assert "test@example.com" not in sanitized_item["output"]
@pytest.mark.asyncio
async def test_realtime_function_call_output_guardrail_allows_clean_output(monkeypatch: pytest.MonkeyPatch):
"""
@ -1708,7 +1719,6 @@ async def test_realtime_function_call_output_guardrail_allows_clean_output(monke
assert len(forwarded) == 1, f"Clean function_call_output should be forwarded, got: {forwarded}"
@pytest.mark.asyncio
async def test_realtime_text_input_guardrail_uses_pre_call_mode(monkeypatch: pytest.MonkeyPatch):
"""
@ -1744,7 +1754,6 @@ async def test_realtime_text_input_guardrail_uses_pre_call_mode(monkeypatch: pyt
)
@pytest.mark.asyncio
async def test_realtime_session_created_injects_session_update_for_audio_guardrail(monkeypatch: pytest.MonkeyPatch):
"""
@ -1801,7 +1810,6 @@ async def test_realtime_session_created_injects_session_update_for_audio_guardra
)
@pytest.mark.asyncio
async def test_realtime_session_created_does_not_inject_session_update_for_pre_call_only(
monkeypatch: pytest.MonkeyPatch,
@ -1846,7 +1854,6 @@ async def test_realtime_session_created_does_not_inject_session_update_for_pre_c
assert len(session_updates) == 0, f"pre_call-only guardrail must not inject session.update, got: {sent_to_backend}"
@pytest.mark.asyncio
async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad(monkeypatch: pytest.MonkeyPatch):
"""Model Armor-style pre_call + post_call must not gate audio VAD."""
@ -1862,17 +1869,17 @@ async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad(monke
litellm,
"callbacks",
[
ModelArmorStyleGuardrail(
guardrail_name="model_armor_all_pre_call",
event_hook=GuardrailEventHooks.pre_call,
default_on=False,
),
ModelArmorStyleGuardrail(
guardrail_name="model_armor_all_post_call",
event_hook=GuardrailEventHooks.post_call,
default_on=False,
),
],
ModelArmorStyleGuardrail(
guardrail_name="model_armor_all_pre_call",
event_hook=GuardrailEventHooks.pre_call,
default_on=False,
),
ModelArmorStyleGuardrail(
guardrail_name="model_armor_all_post_call",
event_hook=GuardrailEventHooks.post_call,
default_on=False,
),
],
)
client_ws = MagicMock()
@ -1896,7 +1903,6 @@ async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad(monke
assert streaming._has_audio_transcription_guardrails() is False
@pytest.mark.asyncio
async def test_end_session_after_n_fails_closes_connection(monkeypatch: pytest.MonkeyPatch):
"""
@ -1943,7 +1949,6 @@ async def test_end_session_after_n_fails_closes_connection(monkeypatch: pytest.M
assert streaming._violation_count == 2
@pytest.mark.asyncio
async def test_on_violation_end_session_closes_on_first_fail(monkeypatch: pytest.MonkeyPatch):
"""
@ -1989,7 +1994,6 @@ async def test_on_violation_end_session_closes_on_first_fail(monkeypatch: pytest
assert streaming._violation_count == 1
@pytest.mark.asyncio
async def test_provider_path_suppresses_duplicate_session_created_after_synthetic():
client_ws = MagicMock()
@ -2952,7 +2956,9 @@ async def test_log_messages_routes_async_logging_through_bounded_worker():
mock_worker.ensure_initialized_and_enqueue.assert_called_once()
enqueued = mock_worker.ensure_initialized_and_enqueue.call_args
assert (enqueued.args or tuple(enqueued.kwargs.values()))[0] is logging_obj.dispatch_success_handlers.return_value
assert (enqueued.args or tuple(enqueued.kwargs.values()))[
0
] is logging_obj.dispatch_success_handlers.return_value
logging_obj.dispatch_success_handlers.assert_called_once_with(streaming.messages, prefer_async_handlers=True)
logging_obj.success_handler.assert_not_called()
# the bare create_task path must no longer be used for success logging

View file

@ -7,6 +7,7 @@ from unittest.mock import MagicMock
import pytest
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.realtime.handler import BedrockRealtime
from litellm.llms.bedrock.realtime.transformation import BedrockRealtimeConfig
@ -55,6 +56,33 @@ class FakeBedrockStream:
class FakeLogging:
def __init__(self, trace_id="trace-nova-sonic"):
self.litellm_trace_id = trace_id
self.dispatched_results = []
self.model_call_details = {}
self.pre_call_args = []
def pre_call(self, input=None, api_key="", model=None, additional_args=None):
self.pre_call_args.append(
{"input": input, "api_key": api_key, "model": model, "additional_args": additional_args or {}}
)
async def dispatch_success_handlers(self, result=None, prefer_async_handlers=False, **kwargs):
self.dispatched_results.append(result)
@pytest.fixture(autouse=True)
def drain_bedrock_realtime_logging_worker(monkeypatch):
pending = []
def capture_enqueue(coro):
pending.append(coro)
monkeypatch.setattr(
"litellm.litellm_core_utils.realtime_streaming.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
capture_enqueue,
)
yield pending
for coro in pending:
coro.close()
class DisconnectingClientWS:
@ -89,6 +117,37 @@ class EndedBedrockStream:
return (None, EndedBedrockReceiver())
class ScriptedBedrockChunk:
def __init__(self, payload):
self.bytes_ = json.dumps(payload).encode("utf-8")
class ScriptedBedrockResult:
def __init__(self, payload):
self.value = ScriptedBedrockChunk(payload)
class ScriptedBedrockReceiver:
def __init__(self, payloads):
self._payloads = list(payloads)
async def receive(self):
if not self._payloads:
return None
return ScriptedBedrockResult(self._payloads.pop(0))
class ScriptedBedrockStream:
"""Replays a fixed list of Bedrock event payloads, then ends the stream."""
def __init__(self, payloads):
self.input_stream = FakeInputStream()
self._receiver = ScriptedBedrockReceiver(payloads)
async def await_output(self):
return (None, self._receiver)
class RealtimeClientWS:
def __init__(self):
self.closed = False
@ -151,7 +210,9 @@ def stub_aws_sdk_client(monkeypatch):
async def invoke_model_with_bidirectional_stream(self, operation_input):
captured["operation_input"] = operation_input
return ImmediatelyEndingBedrockStream()
# Tests that need the session to see Bedrock frames set captured["stream_events"];
# with none set this replays nothing and ends immediately.
return ScriptedBedrockStream(captured.get("stream_events", []))
package = types.ModuleType("aws_sdk_bedrock_runtime")
client_module = types.ModuleType("aws_sdk_bedrock_runtime.client")
@ -311,6 +372,161 @@ class TestBedrockRealtimeSessionLifecycle:
assert first_event["session"]["id"] == "trace-nova-sonic"
assert first_event["session"]["model"] == "amazon.nova-sonic-v1:0"
@pytest.mark.asyncio
async def test_session_dispatches_logged_events_via_realtime_streaming(
self, stub_aws_sdk_client, drain_bedrock_realtime_logging_worker
):
handler = BedrockRealtime()
websocket = RealtimeClientWS()
logging_obj = FakeLogging()
await handler.async_realtime(
model="amazon.nova-sonic-v1:0",
websocket=websocket,
logging_obj=logging_obj,
aws_region_name="us-east-1",
aws_access_key_id="k",
aws_secret_access_key="s",
)
assert logging_obj.pre_call_args
assert len(drain_bedrock_realtime_logging_worker) == 1
await drain_bedrock_realtime_logging_worker.pop()
assert logging_obj.dispatched_results
dispatched = logging_obj.dispatched_results[0]
assert any(event.get("type") == "session.created" for event in dispatched)
@pytest.mark.asyncio
async def test_unbilled_usage_at_session_close_is_flushed_into_the_spend_log(
self, stub_aws_sdk_client, drain_bedrock_realtime_logging_worker
):
"""
Usage that arrives while a response is open is not billed by any response.done during
the session, so closing the session must flush it or the turn is never charged.
"""
stub_aws_sdk_client["stream_events"] = [
{"event": {"contentStart": {"role": "ASSISTANT", "type": "TEXT"}}},
{
"event": {
"usageEvent": {
"totalInputTokens": 12,
"totalOutputTokens": 23,
"totalTokens": 35,
"details": {
"total": {
"input": {"speechTokens": 10, "textTokens": 2},
"output": {"speechTokens": 20, "textTokens": 3},
}
},
}
}
},
]
handler = BedrockRealtime()
logging_obj = FakeLogging()
await handler.async_realtime(
model="amazon.nova-sonic-v1:0",
websocket=RealtimeClientWS(),
logging_obj=logging_obj,
aws_region_name="us-east-1",
aws_access_key_id="k",
aws_secret_access_key="s",
)
await drain_bedrock_realtime_logging_worker.pop()
dispatched = logging_obj.dispatched_results[0]
done_events = [event for event in dispatched if event.get("type") == "response.done"]
assert len(done_events) == 1, "session close did not flush the open turn's usage"
usage = done_events[0]["response"]["usage"]
assert usage["input_tokens"] == 12
assert usage["output_tokens"] == 23
assert usage["total_tokens"] == 35
@pytest.mark.asyncio
async def test_client_session_update_reaches_spend_logging(self, stub_aws_models):
"""Declared tools and instructions must reach the spend log via store_input."""
handler = BedrockRealtime()
client_ws = DisconnectingClientWS(
[
json.dumps(
{
"type": "session.update",
"session": {
"instructions": "be brief",
"tools": [{"type": "function", "name": "get_weather"}],
},
}
)
]
)
realtime_streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=None,
logging_obj=FakeLogging(),
model="amazon.nova-sonic-v1:0",
)
await handler._forward_client_to_bedrock(
client_ws,
FakeBedrockStream(),
BedrockRealtimeConfig(),
"amazon.nova-sonic-v1:0",
{},
FakeLogging(),
realtime_streaming,
)
assert realtime_streaming.session_tools == [{"type": "function", "name": "get_weather"}]
assert {"role": "system", "content": "be brief"} in realtime_streaming.input_messages
@pytest.mark.asyncio
async def test_tool_call_reaches_spend_logging_via_response_done(self):
"""
Bedrock tool calls must be billable through the shared RealTimeStreaming collector,
which reads function_call items off response.done, with no Bedrock-specific plumbing.
"""
handler = BedrockRealtime()
client_ws = RealtimeClientWS()
logging_obj = FakeLogging()
realtime_streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=None,
logging_obj=logging_obj,
model="amazon.nova-2-sonic-v1:0",
)
await handler._forward_bedrock_to_client(
ScriptedBedrockStream(
[
{"event": {"contentStart": {"role": "TOOL", "type": "TOOL"}}},
{
"event": {
"toolUse": {
"toolUseId": "tool_call_1",
"toolName": "get_weather",
"content": json.dumps({"location": "Seattle"}),
}
}
},
]
),
client_ws,
BedrockRealtimeConfig(),
"amazon.nova-2-sonic-v1:0",
logging_obj,
{},
realtime_streaming,
)
assert realtime_streaming.tool_calls == [
{
"id": "tool_call_1",
"type": "function",
"function": {"name": "get_weather", "arguments": json.dumps({"location": "Seattle"})},
}
]
@pytest.mark.asyncio
async def test_session_update_is_acked_with_session_updated(self, stub_aws_models):
handler = BedrockRealtime()
@ -320,9 +536,7 @@ class TestBedrockRealtimeSessionLifecycle:
[json.dumps({"type": "session.update", "session": {"instructions": "hi", "modalities": ["text"]}})]
)
await handler._forward_client_to_bedrock(
client_ws, stream, config, "amazon.nova-sonic-v1:0", {}, FakeLogging()
)
await handler._forward_client_to_bedrock(client_ws, stream, config, "amazon.nova-sonic-v1:0", {}, FakeLogging())
acked = [json.loads(message) for message in client_ws.sent_to_client]
updated = [event for event in acked if event["type"] == "session.updated"]
@ -334,9 +548,7 @@ class TestBedrockRealtimeSessionLifecycle:
handler = BedrockRealtime()
config = BedrockRealtimeConfig()
stream = FakeBedrockStream()
client_ws = DisconnectingClientWS(
[json.dumps({"type": "session.update", "session": {"instructions": "hi"}})]
)
client_ws = DisconnectingClientWS([json.dumps({"type": "session.update", "session": {"instructions": "hi"}})])
await handler._forward_client_to_bedrock(client_ws, stream, config, "amazon.nova-sonic-v1:0", {})

View file

@ -12,7 +12,34 @@ from litellm.llms.bedrock.realtime.transformation import (
BedrockRealtimeConfig,
)
from litellm.llms.bedrock.realtime.trigger_audio import ready_trigger_pcm
from litellm.types.llms.openai import OpenAIRealtimeEventTypes
# The OpenAI realtime function-call lifecycle a single Nova Sonic toolUse must expand into,
# when an assistant block already opened the response the tool call belongs to.
_TOOL_CALL_EVENT_SEQUENCE = [
"response.output_item.added",
"conversation.item.added",
"response.function_call_arguments.delta",
"response.function_call_arguments.done",
"response.output_item.done",
"response.done",
]
# Nova Sonic opens tool turns with contentStart role TOOL, which opens no response, so the
# tool call has to open one itself before it can close it.
_TOOL_CALL_EVENT_SEQUENCE_NEW_RESPONSE = ["response.created", *_TOOL_CALL_EVENT_SEQUENCE]
def _only(events, event_type):
matches = [event for event in events if event["type"] == event_type]
assert len(matches) == 1, f"expected exactly one {event_type}, got {len(matches)}"
return matches[0]
def _response_id_of(event):
"""The response id an event is bound to, or None for events that carry no response id."""
if event["type"] in ("response.created", "response.done"):
return event["response"]["id"]
return event.get("response_id")
class TestBedrockRealtimeConfig:
@ -557,16 +584,322 @@ class TestBedrockRealtimeResponseTransformation:
},
)
# Check for function call event
assert len(result["response"]) == 1
function_call = result["response"][0]
assert function_call["type"] == "response.function_call_arguments.done"
assert [msg["type"] for msg in result["response"]] == _TOOL_CALL_EVENT_SEQUENCE
function_call = _only(result["response"], "response.function_call_arguments.done")
assert function_call["call_id"] == "tool_call_123"
assert function_call["name"] == "get_weather"
assert json.loads(function_call["arguments"]) == {"location": "San Francisco"}
# Verify arguments are properly formatted
args = json.loads(function_call["arguments"])
assert args["location"] == "San Francisco"
def test_transform_tool_use_response_with_content_field(self):
"""Test toolUse response transformation with Nova 2 Sonic `content` field"""
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
tool_use_message = {
"event": {
"toolUse": {
"toolUseId": "tool_call_123",
"toolName": "get_weather",
"content": json.dumps({"location": "San Francisco"}),
}
}
}
result = config.transform_realtime_response(
json.dumps(tool_use_message),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input={
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_123",
"current_response_id": "resp_123",
"current_conversation_id": "conv_123",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": "text",
},
)
assert [msg["type"] for msg in result["response"]] == _TOOL_CALL_EVENT_SEQUENCE
function_call = _only(result["response"], "response.function_call_arguments.done")
assert function_call["call_id"] == "tool_call_123"
assert function_call["name"] == "get_weather"
assert json.loads(function_call["arguments"]) == {"location": "San Francisco"}
done = _only(result["response"], "response.done")
assert done["response"]["id"] == "resp_123"
assert done["response"]["output"] == [
{
"id": "item_123",
"object": "realtime.item",
"type": "function_call",
"status": "completed",
"call_id": "tool_call_123",
"name": "get_weather",
"arguments": function_call["arguments"],
}
]
assert result["current_response_id"] is None
assert result["current_output_item_id"] is None
def test_transform_tool_use_event_directly(self):
"""transform_tool_use_event emits the full OpenAI function-call lifecycle"""
config = BedrockRealtimeConfig()
# Missing IDs are minted (Nova Sonic starts tool turns with contentStart role=TOOL)
events, tool_call_id, tool_name = config.transform_tool_use_event(
{
"toolUse": {
"toolUseId": "tool_call_no_ids",
"toolName": "get_weather",
"content": json.dumps({"location": "Seattle"}),
}
},
None,
None,
"conv_1",
)
assert [event["type"] for event in events] == _TOOL_CALL_EVENT_SEQUENCE_NEW_RESPONSE
function_call = _only(events, "response.function_call_arguments.done")
assert _only(events, "response.created")["response"]["id"] == function_call["response_id"]
assert _only(events, "response.created")["response"]["conversation_id"] == "conv_1"
assert function_call["call_id"] == "tool_call_no_ids"
assert function_call["name"] == "get_weather"
assert function_call["response_id"].startswith("resp_")
assert function_call["item_id"].startswith("item_")
assert json.loads(function_call["arguments"]) == {"location": "Seattle"}
assert tool_call_id == "tool_call_no_ids"
assert tool_name == "get_weather"
# Every event in the turn shares the minted response/item ids
assert {_response_id_of(event) for event in events} - {None} == {function_call["response_id"]}
assert _only(events, "response.output_item.added")["item"]["id"] == function_call["item_id"]
assert _only(events, "response.output_item.done")["item"]["id"] == function_call["item_id"]
# The added item is in_progress with empty args; the done item carries the parsed args
assert _only(events, "response.output_item.added")["item"]["status"] == "in_progress"
assert _only(events, "response.output_item.added")["item"]["arguments"] == ""
assert _only(events, "conversation.item.added")["item"]["arguments"] == ""
assert _only(events, "response.function_call_arguments.delta")["delta"] == function_call["arguments"]
assert _only(events, "response.output_item.done")["item"]["status"] == "completed"
assert _only(events, "response.output_item.done")["item"]["arguments"] == function_call["arguments"]
# response.done closes the turn and carries the call so spend logging can harvest it
done = _only(events, "response.done")
assert done["response"]["id"] == function_call["response_id"]
assert done["response"]["conversation_id"] == "conv_1"
assert done["response"]["status"] == "completed"
assert done["response"]["output"][0]["call_id"] == "tool_call_no_ids"
assert done["response"]["output"][0]["type"] == "function_call"
# Explicit ids are reused rather than minted
events, _, _ = config.transform_tool_use_event(
{
"toolUse": {
"toolUseId": "tool_call_123",
"toolName": "get_weather",
"content": json.dumps({"location": "San Francisco"}),
}
},
"item_123",
"resp_123",
"conv_1",
)
function_call = _only(events, "response.function_call_arguments.done")
assert function_call["response_id"] == "resp_123"
assert function_call["item_id"] == "item_123"
assert json.loads(function_call["arguments"]) == {"location": "San Francisco"}
# Legacy `input` field is still honoured when `content` is absent
events, _, _ = config.transform_tool_use_event(
{
"toolUse": {
"toolUseId": "tool_call_legacy",
"toolName": "get_weather",
"input": json.dumps({"location": "Boston"}),
}
},
"item_123",
"resp_123",
"conv_1",
)
assert json.loads(_only(events, "response.function_call_arguments.done")["arguments"]) == {"location": "Boston"}
# Invalid JSON content falls back to empty arguments
events, _, _ = config.transform_tool_use_event(
{
"toolUse": {
"toolUseId": "tool_call_124",
"toolName": "get_weather",
"content": "not valid json",
}
},
"item_123",
"resp_123",
"conv_1",
)
assert json.loads(_only(events, "response.function_call_arguments.done")["arguments"]) == {}
def test_transform_realtime_response_persists_minted_tool_ids(self):
"""TOOL-first turns must write minted response/item ids into session state"""
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
state = {
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": None,
"current_response_id": None,
"current_conversation_id": "conv_123",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": None,
}
content_start_result = config.transform_realtime_response(
json.dumps({"event": {"contentStart": {"role": "TOOL", "type": "TOOL"}}}),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
assert content_start_result["response"] == []
assert content_start_result["current_delta_type"] is None
state.update(
{
"current_output_item_id": content_start_result["current_output_item_id"],
"current_response_id": content_start_result["current_response_id"],
"current_conversation_id": content_start_result["current_conversation_id"],
"current_delta_chunks": content_start_result["current_delta_chunks"],
"current_item_chunks": content_start_result["current_item_chunks"],
"current_delta_type": content_start_result["current_delta_type"],
}
)
tool_use_message = {
"event": {
"toolUse": {
"toolUseId": "tool_call_state",
"toolName": "get_weather",
"content": json.dumps({"location": "Seattle"}),
}
}
}
result = config.transform_realtime_response(
json.dumps(tool_use_message),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
assert [msg["type"] for msg in result["response"]] == _TOOL_CALL_EVENT_SEQUENCE_NEW_RESPONSE
function_call = _only(result["response"], "response.function_call_arguments.done")
assert _only(result["response"], "response.created")["response"]["id"] == function_call["response_id"]
assert function_call["response_id"].startswith("resp_")
assert function_call["item_id"].startswith("item_")
assert json.loads(function_call["arguments"]) == {"location": "Seattle"}
# The tool turn closes the response it minted, so no in-progress response is orphaned
# and the ids cannot leak into the post-tool assistant turn.
tool_done = _only(result["response"], "response.done")
assert tool_done["response"]["id"] == function_call["response_id"]
assert result["current_response_id"] is None
assert result["current_output_item_id"] is None
content_end_message = {
"event": {
"contentEnd": {
"stopReason": "TOOL_USE",
"type": "TOOL",
}
}
}
follow_up = config.transform_realtime_response(
json.dumps(content_end_message),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input={
"session_configuration_request": result["session_configuration_request"],
"current_output_item_id": result["current_output_item_id"],
"current_response_id": result["current_response_id"],
"current_conversation_id": result["current_conversation_id"],
"current_delta_chunks": result["current_delta_chunks"],
"current_item_chunks": result["current_item_chunks"],
"current_delta_type": result["current_delta_type"],
},
)
assert follow_up["current_response_id"] is None
assert follow_up["current_output_item_id"] is None
assert follow_up["current_delta_type"] is None
# The tool turn already emitted response.done; TOOL contentEnd must not emit a second
# one, nor an unpaired message-shaped output_item.done.
assert follow_up["response"] == []
post_tool_state = {
"session_configuration_request": follow_up["session_configuration_request"],
"current_output_item_id": follow_up["current_output_item_id"],
"current_response_id": follow_up["current_response_id"],
"current_conversation_id": follow_up["current_conversation_id"],
"current_delta_chunks": follow_up["current_delta_chunks"],
"current_item_chunks": follow_up["current_item_chunks"],
"current_delta_type": follow_up["current_delta_type"],
}
assistant_start = config.transform_realtime_response(
json.dumps({"event": {"contentStart": {"role": "ASSISTANT", "type": "TEXT"}}}),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input=post_tool_state,
)
assert assistant_start["current_response_id"] is not None
assert assistant_start["current_output_item_id"] is not None
assert assistant_start["current_response_id"] != function_call["response_id"]
assert assistant_start["current_output_item_id"] != function_call["item_id"]
created = [msg for msg in assistant_start["response"] if msg["type"] == "response.created"][0]
added = [msg for msg in assistant_start["response"] if msg["type"] == "response.output_item.added"][0]
assert created["response"]["id"] == assistant_start["current_response_id"]
assert added["item"]["id"] == assistant_start["current_output_item_id"]
assert created["response"]["id"] != function_call["response_id"]
assert added["item"]["id"] != function_call["item_id"]
def test_tool_content_end_does_not_emit_message_output_item_done(self):
"""
TOOL contentEnd must stay silent: the tool turn already closed its own response, and a
message-shaped output_item.done here would have no matching output_item.added.
"""
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
content_end_message = {
"event": {
"contentEnd": {
"stopReason": "TOOL_USE",
"type": "TOOL",
}
}
}
result = config.transform_realtime_response(
json.dumps(content_end_message),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input={
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_open_assistant_turn",
"current_response_id": "resp_open_assistant_turn",
"current_conversation_id": "conv_123",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": "text",
},
)
assert result["response"] == []
assert result["current_output_item_id"] is None
assert result["current_delta_type"] is None
# A TOOL block that produced no toolUse leaves the assistant response open rather than
# dropping its id, so the next assistant block reuses it instead of orphaning it.
assert result["current_response_id"] == "resp_open_assistant_turn"
def test_transform_content_end_text(self):
"""Test contentEnd for text response"""
@ -827,5 +1160,571 @@ class TestBedrockRealtimeSessionEvents:
assert event["session"]["modalities"] == ["text", "audio"]
class TestBedrockRealtimeContentBlockLifecycle:
"""
Bedrock streams discrete content blocks. Session state must follow block
boundaries so text/audio/tool blocks cannot leak into each other.
"""
def _state(self, **overrides):
base = {
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": None,
"current_response_id": None,
"current_conversation_id": "conv_1",
"current_delta_chunks": None,
"current_item_chunks": [],
"current_delta_type": None,
}
base.update(overrides)
return base
def _apply(self, config, logging_obj, state, message):
result = config.transform_realtime_response(
json.dumps(message),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
state.update(
{
"current_output_item_id": result["current_output_item_id"],
"current_response_id": result["current_response_id"],
"current_conversation_id": result["current_conversation_id"],
"current_delta_chunks": result["current_delta_chunks"],
"current_item_chunks": result["current_item_chunks"],
"current_delta_type": result["current_delta_type"],
}
)
return result
def test_tool_block_does_not_leak_prior_text_into_next_assistant_turn(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
state = self._state()
self._apply(config, logging_obj, state, {"event": {"contentStart": {"role": "ASSISTANT", "type": "TEXT"}}})
first_response_id = state["current_response_id"]
self._apply(
config,
logging_obj,
state,
{"event": {"textOutput": {"content": "I will check the weather."}}},
)
assert state["current_delta_chunks"] is not None
assert len(state["current_delta_chunks"]) == 1
self._apply(
config,
logging_obj,
state,
{"event": {"contentEnd": {"stopReason": "PARTIAL_TURN", "type": "TEXT"}}},
)
assert state["current_delta_chunks"] is None
assert state["current_delta_type"] is None
assert state["current_output_item_id"] is None
assert state["current_response_id"] == first_response_id
self._apply(config, logging_obj, state, {"event": {"contentStart": {"role": "TOOL", "type": "TOOL"}}})
assert state["current_delta_chunks"] is None
assert state["current_response_id"] == first_response_id
tool_result = self._apply(
config,
logging_obj,
state,
{
"event": {
"toolUse": {
"toolUseId": "tool_1",
"toolName": "get_weather",
"content": json.dumps({"location": "Seattle"}),
}
}
},
)
assert [msg["type"] for msg in tool_result["response"]] == _TOOL_CALL_EVENT_SEQUENCE
# The response the assistant text block opened is closed by the tool turn instead of
# being left in_progress forever once the ids are cleared.
tool_done = _only(tool_result["response"], "response.done")
assert tool_done["response"]["id"] == first_response_id
assert tool_done["response"]["output"][0]["call_id"] == "tool_1"
assert state["current_response_id"] is None
assert state["current_output_item_id"] is None
assert state["current_delta_chunks"] is None
assert state["current_delta_type"] is None
tool_content_end = self._apply(
config,
logging_obj,
state,
{"event": {"contentEnd": {"stopReason": "TOOL_USE", "type": "TOOL"}}},
)
assert tool_content_end["response"] == []
assert state["current_response_id"] is None
assert state["current_output_item_id"] is None
post_tool = self._apply(
config,
logging_obj,
state,
{"event": {"contentStart": {"role": "ASSISTANT", "type": "TEXT"}}},
)
assert state["current_response_id"] != first_response_id
assert state["current_delta_chunks"] is None
assert [msg["type"] for msg in post_tool["response"]].count("response.created") == 1
self._apply(
config,
logging_obj,
state,
{"event": {"textOutput": {"content": "It is sunny in Seattle."}}},
)
done = self._apply(
config,
logging_obj,
state,
{"event": {"contentEnd": {"stopReason": "END_TURN", "type": "TEXT"}}},
)
text_done = [msg for msg in done["response"] if msg["type"] == "response.text.done"][0]
assert text_done["text"] == "It is sunny in Seattle."
assert "I will check the weather." not in text_done["text"]
assert any(msg["type"] == "response.done" for msg in done["response"])
assert state["current_response_id"] is None
assert state["current_delta_chunks"] is None
def test_every_created_response_is_closed_across_a_tool_turn(self):
"""
Realtime clients track in-progress responses by id. A tool turn that drops the
response id without a matching response.done leaves one open forever.
"""
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
state = self._state()
turn = [
{"event": {"contentStart": {"role": "ASSISTANT", "type": "TEXT"}}},
{"event": {"textOutput": {"content": "Let me check."}}},
{"event": {"contentEnd": {"stopReason": "PARTIAL_TURN", "type": "TEXT"}}},
{"event": {"contentStart": {"role": "TOOL", "type": "TOOL"}}},
{
"event": {
"toolUse": {
"toolUseId": "tool_1",
"toolName": "get_weather",
"content": json.dumps({"location": "Seattle"}),
}
}
},
{"event": {"contentEnd": {"stopReason": "TOOL_USE", "type": "TOOL"}}},
{"event": {"contentStart": {"role": "ASSISTANT", "type": "TEXT"}}},
{"event": {"textOutput": {"content": "It is sunny."}}},
{"event": {"contentEnd": {"stopReason": "END_TURN", "type": "TEXT"}}},
]
emitted = [msg for event in turn for msg in self._apply(config, logging_obj, state, event)["response"]]
created = [msg["response"]["id"] for msg in emitted if msg["type"] == "response.created"]
done = [msg["response"]["id"] for msg in emitted if msg["type"] == "response.done"]
assert len(created) == 2
assert created == done
assert state["current_response_id"] is None
def test_every_response_is_opened_and_closed_on_a_tool_first_turn(self):
"""
Nova Sonic opens tool turns with contentStart role TOOL, which emits nothing, so the
tool call is the first thing in the session. It has to open the response it closes,
or the client sees a response.done for an id it never saw created.
"""
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
state = self._state(current_conversation_id=None)
turn = [
{"event": {"contentStart": {"role": "TOOL", "type": "TOOL"}}},
{
"event": {
"toolUse": {
"toolUseId": "tool_1",
"toolName": "get_weather",
"content": json.dumps({"location": "Seattle"}),
}
}
},
{"event": {"contentEnd": {"stopReason": "TOOL_USE", "type": "TOOL"}}},
]
emitted = [msg for event in turn for msg in self._apply(config, logging_obj, state, event)["response"]]
created = [msg["response"]["id"] for msg in emitted if msg["type"] == "response.created"]
done = [msg["response"]["id"] for msg in emitted if msg["type"] == "response.done"]
assert len(created) == 1
assert created == done
# Every event in the turn is bound to that one response.
assert {_response_id_of(msg) for msg in emitted} - {None} == set(created)
assert state["current_response_id"] is None
def test_second_assistant_content_block_reuses_response_not_item(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
state = self._state()
first = self._apply(
config, logging_obj, state, {"event": {"contentStart": {"role": "ASSISTANT", "type": "TEXT"}}}
)
response_id = state["current_response_id"]
first_item = state["current_output_item_id"]
assert sum(1 for msg in first["response"] if msg["type"] == "response.created") == 1
self._apply(
config,
logging_obj,
state,
{"event": {"contentEnd": {"stopReason": "PARTIAL_TURN", "type": "TEXT"}}},
)
second = self._apply(
config, logging_obj, state, {"event": {"contentStart": {"role": "ASSISTANT", "type": "AUDIO"}}}
)
assert state["current_response_id"] == response_id
assert state["current_output_item_id"] != first_item
assert sum(1 for msg in second["response"] if msg["type"] == "response.created") == 0
assert sum(1 for msg in second["response"] if msg["type"] == "response.output_item.added") == 1
class TestBedrockRealtimeToolArgumentParsing:
"""Nova Sonic tool args arrive in several shapes; none may crash or leak a raw wire value."""
def _args(self, tool_use: dict) -> str:
config = BedrockRealtimeConfig()
events, _, _ = config.transform_tool_use_event({"toolUse": tool_use}, "item_1", "resp_1", "conv_1")
return _only(events, "response.function_call_arguments.done")["arguments"]
def test_already_decoded_object_content_is_passed_through(self):
assert json.loads(self._args({"toolUseId": "t", "toolName": "f", "content": {"location": "Seattle"}})) == {
"location": "Seattle"
}
def test_empty_content_falls_back_to_empty_arguments(self):
assert json.loads(self._args({"toolUseId": "t", "toolName": "f", "content": ""})) == {}
def test_missing_content_and_input_yields_empty_arguments(self):
assert json.loads(self._args({"toolUseId": "t", "toolName": "f"})) == {}
def test_non_json_content_yields_empty_arguments(self):
assert json.loads(self._args({"toolUseId": "t", "toolName": "f", "content": "not json"})) == {}
class TestBedrockRealtimeUsageAccounting:
def _usage_event(
self,
*,
input_speech: int,
input_text: int,
output_speech: int,
output_text: int,
total_input: int | None = None,
total_output: int | None = None,
total: int | None = None,
) -> dict:
resolved_input = total_input if total_input is not None else input_speech + input_text
resolved_output = total_output if total_output is not None else output_speech + output_text
resolved_total = total if total is not None else resolved_input + resolved_output
return {
"event": {
"usageEvent": {
"completionId": "completion_1",
"details": {
"total": {
"input": {"speechTokens": input_speech, "textTokens": input_text},
"output": {"speechTokens": output_speech, "textTokens": output_text},
}
},
"promptName": "prompt_1",
"sessionId": "session_1",
"totalInputTokens": resolved_input,
"totalOutputTokens": resolved_output,
"totalTokens": resolved_total,
}
}
}
def test_usage_event_fills_response_done_turn_delta(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
state = {
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_1",
"current_response_id": "resp_1",
"current_conversation_id": "conv_1",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": "audio",
}
config.transform_realtime_response(
json.dumps(self._usage_event(input_speech=10, input_text=2, output_speech=20, output_text=3)),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
first_done = config.transform_realtime_response(
json.dumps({"event": {"contentEnd": {"stopReason": "END_TURN", "type": "AUDIO"}}}),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
done_events = [msg for msg in first_done["response"] if msg["type"] == "response.done"]
assert len(done_events) == 1
usage = done_events[0]["response"]["usage"]
assert usage["input_tokens"] == 12
assert usage["output_tokens"] == 23
assert usage["total_tokens"] == 35
assert usage["input_token_details"]["audio_tokens"] == 10
assert usage["input_token_details"]["text_tokens"] == 2
assert usage["output_token_details"]["audio_tokens"] == 20
assert usage["output_token_details"]["text_tokens"] == 3
config.transform_realtime_response(
json.dumps(
self._usage_event(
input_speech=15,
input_text=2,
output_speech=30,
output_text=3,
total_input=17,
total_output=33,
total=50,
)
),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input={
**state,
"current_output_item_id": "item_2",
"current_response_id": "resp_2",
},
)
second_done = config.transform_realtime_response(
json.dumps({"event": {"contentEnd": {"stopReason": "END_TURN", "type": "AUDIO"}}}),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input={
**state,
"current_output_item_id": "item_2",
"current_response_id": "resp_2",
},
)
second_usage = [msg for msg in second_done["response"] if msg["type"] == "response.done"][0]["response"][
"usage"
]
assert second_usage["input_tokens"] == 5
assert second_usage["output_tokens"] == 10
assert second_usage["total_tokens"] == 15
assert second_usage["input_token_details"]["audio_tokens"] == 5
assert second_usage["output_token_details"]["audio_tokens"] == 10
def test_late_usage_event_after_response_id_cleared_emits_response_done(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_late_usage"
state = {
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_1",
"current_response_id": "resp_1",
"current_conversation_id": "conv_1",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": "audio",
}
end_turn = config.transform_realtime_response(
json.dumps({"event": {"contentEnd": {"stopReason": "END_TURN", "type": "AUDIO"}}}),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
assert any(msg["type"] == "response.done" for msg in end_turn["response"])
assert end_turn["current_response_id"] is None
state["current_response_id"] = end_turn["current_response_id"]
state["current_output_item_id"] = end_turn["current_output_item_id"]
state["current_conversation_id"] = end_turn["current_conversation_id"]
late_usage = config.transform_realtime_response(
json.dumps(self._usage_event(input_speech=10, input_text=2, output_speech=20, output_text=3)),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
done_events = [msg for msg in late_usage["response"] if msg["type"] == "response.done"]
assert len(done_events) == 1
usage = done_events[0]["response"]["usage"]
assert usage["input_tokens"] == 12
assert usage["output_tokens"] == 23
assert usage["total_tokens"] == 35
assert not config.has_unbilled_usage()
def test_tool_content_end_with_end_turn_stop_reason_emits_response_done(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_tool_end_turn"
state = {
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_tool",
"current_response_id": "resp_tool",
"current_conversation_id": "conv_tool",
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": None,
}
config.transform_realtime_response(
json.dumps(self._usage_event(input_speech=4, input_text=1, output_speech=0, output_text=2)),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
tool_end = config.transform_realtime_response(
json.dumps({"event": {"contentEnd": {"stopReason": "END_TURN", "type": "TOOL"}}}),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
done_events = [msg for msg in tool_end["response"] if msg["type"] == "response.done"]
assert len(done_events) == 1
assert done_events[0]["response"]["id"] == "resp_tool"
assert done_events[0]["response"]["usage"]["input_tokens"] == 5
assert done_events[0]["response"]["usage"]["output_tokens"] == 2
assert tool_end["current_response_id"] is None
assert not config.has_unbilled_usage()
def test_tool_turn_does_not_double_bill_across_late_usage_and_completion_end(self):
"""The tool response.done, a later usageEvent, and completionEnd must each bill once."""
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_tool_usage"
state = {
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": "item_1",
"current_response_id": "resp_1",
"current_conversation_id": "conv_1",
"current_delta_chunks": None,
"current_item_chunks": [],
"current_delta_type": None,
}
config.transform_realtime_response(
json.dumps(self._usage_event(input_speech=4, input_text=0, output_speech=6, output_text=0)),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
tool_result = config.transform_realtime_response(
json.dumps(
{
"event": {
"toolUse": {
"toolUseId": "tool_1",
"toolName": "get_weather",
"content": json.dumps({"location": "Seattle"}),
}
}
}
),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
tool_usage = _only(tool_result["response"], "response.done")["response"]["usage"]
assert tool_usage["input_tokens"] == 4
assert tool_usage["output_tokens"] == 6
assert not config.has_unbilled_usage()
state["current_response_id"] = tool_result["current_response_id"]
state["current_output_item_id"] = tool_result["current_output_item_id"]
# Cumulative usage grows after the tool turn; only the delta may be billed again.
late = config.transform_realtime_response(
json.dumps(self._usage_event(input_speech=10, input_text=0, output_speech=15, output_text=0)),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
late_usage = _only(late["response"], "response.done")["response"]["usage"]
assert late_usage["input_tokens"] == 6
assert late_usage["output_tokens"] == 9
state["current_response_id"] = late["current_response_id"]
completion_end = config.transform_realtime_response(
json.dumps({"event": {"completionEnd": {}}}),
"amazon.nova-2-sonic-v1:0",
logging_obj,
realtime_response_transform_input=state,
)
assert completion_end["response"] == []
assert not config.has_unbilled_usage()
assert config.flush_pending_usage_as_response_done(None, None) == []
def test_non_numeric_token_counts_are_billed_as_zero(self):
"""A malformed usageEvent must not crash the session or bill a bogus amount."""
config = BedrockRealtimeConfig()
config.record_usage_event(
{
"totalInputTokens": "twelve",
"totalOutputTokens": True,
"totalTokens": None,
"details": {"total": {"input": {"speechTokens": None}, "output": {"textTokens": "x"}}},
}
)
assert not config.has_unbilled_usage()
assert config.flush_pending_usage_as_response_done(None, None) == []
def test_flush_pending_usage_on_session_close(self):
config = BedrockRealtimeConfig()
config.record_usage_event(
self._usage_event(input_speech=8, input_text=1, output_speech=16, output_text=2)["event"]["usageEvent"]
)
assert config.has_unbilled_usage()
flushed = config.flush_pending_usage_as_response_done(None, None)
assert len(flushed) == 1
assert flushed[0]["type"] == "response.done"
usage = flushed[0]["response"]["usage"]
assert usage["input_tokens"] == 9
assert usage["output_tokens"] == 18
assert usage["total_tokens"] == 27
assert not config.has_unbilled_usage()
assert config.flush_pending_usage_as_response_done(None, None) == []
def test_completion_end_with_unbilled_usage_mints_response_done(self):
config = BedrockRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_completion_end"
config.record_usage_event(
self._usage_event(input_speech=3, input_text=0, output_speech=6, output_text=0)["event"]["usageEvent"]
)
result = config.transform_realtime_response(
json.dumps({"event": {"completionEnd": {}}}),
"amazon.nova-sonic-v1:0",
logging_obj,
realtime_response_transform_input={
"session_configuration_request": json.dumps({"configured": True}),
"current_output_item_id": None,
"current_response_id": None,
"current_conversation_id": None,
"current_delta_chunks": None,
"current_item_chunks": None,
"current_delta_type": None,
},
)
done_events = [msg for msg in result["response"] if msg["type"] == "response.done"]
assert len(done_events) == 1
assert done_events[0]["response"]["usage"]["input_tokens"] == 3
assert done_events[0]["response"]["usage"]["output_tokens"] == 6
assert not config.has_unbilled_usage()
if __name__ == "__main__":
pytest.main([__file__, "-v"])