diff --git a/Dockerfile b/Dockerfile index 700b0d6525e..356d0297c4b 100644 --- a/Dockerfile +++ b/Dockerfile @@ -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 \ diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index f0d6d02fccf..81f0714a0a3 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -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 \ diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 9125ed6e70a..0336ee4d198 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -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, diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index 3eeb3cb9fc6..8dcc1d19ffd 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -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]) diff --git a/litellm/llms/bedrock/realtime/transformation.py b/litellm/llms/bedrock/realtime/transformation.py index 951bf636b2f..642a076e516 100644 --- a/litellm/llms/bedrock/realtime/transformation.py +++ b/litellm/llms/bedrock/realtime/transformation.py @@ -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, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 7f5e41d45ce..529827fe58f 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index d4b9f4e8cce..157bb45bc2b 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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 = ( diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 45f6b5c55a9..be87bd69150 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -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 ) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 7f5e41d45ce..529827fe58f 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 52e88db753a..263f9f6e2b0 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -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 diff --git a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py index 9efcee192b1..35c4caf51cf 100644 --- a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py +++ b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py @@ -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", {}) diff --git a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_transformation.py b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_transformation.py index ae6b1febd6b..2d3c0360faa 100644 --- a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_transformation.py +++ b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_transformation.py @@ -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"])