diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 4b774269cc2..3dfbdfdc8a8 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -59,7 +59,7 @@ 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 = "", @@ -70,7 +70,7 @@ class RealTimeStreaming: 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 = {} @@ -162,6 +162,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 6ff26b7d3b0..8dcc1d19ffd 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -11,7 +11,7 @@ import contextlib import json from collections.abc import Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Any, Final from pydantic import TypeAdapter @@ -151,7 +151,7 @@ class BedrockRealtime(BaseAWSLLM): ) realtime_streaming: Final = RealTimeStreaming( websocket=websocket, - backend_ws=cast(Any, object()), # cast-ok: Bedrock uses AWS SDK stream; backend_ws unused for store/log + 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, @@ -228,24 +228,6 @@ class BedrockRealtime(BaseAWSLLM): pass raise - @staticmethod - def _collect_tool_call_from_function_call_event( - realtime_streaming: RealTimeStreaming, - message: object, - ) -> None: - if not isinstance(message, dict) or message.get("type") != "response.function_call_arguments.done": - return - realtime_streaming.tool_calls.append( - { # mutable-ok: spend logger tool_calls is a mutable list of JSON dicts - "id": message.get("call_id", ""), - "type": "function", - "function": { # mutable-ok: nested tool function payload - "name": message.get("name", ""), - "arguments": message.get("arguments", "{}"), - }, - } - ) - async def _forward_client_to_bedrock( self, client_ws: Any, @@ -378,7 +360,6 @@ class BedrockRealtime(BaseAWSLLM): for openai_message in openai_messages: if realtime_streaming is not None: realtime_streaming.store_message(openai_message) - self._collect_tool_call_from_function_call_event(realtime_streaming, 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 f3d55a30067..d0936231d52 100644 --- a/litellm/llms/bedrock/realtime/transformation.py +++ b/litellm/llms/bedrock/realtime/transformation.py @@ -9,7 +9,7 @@ import json import uuid as uuid_lib from collections.abc import Mapping from types import MappingProxyType -from typing import Any, Final +from typing import Final from pydantic import BaseModel @@ -20,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, @@ -29,6 +32,7 @@ from litellm.types.llms.openai import ( OpenAIRealtimeResponseDoneObject, OpenAIRealtimeResponseTextDone, OpenAIRealtimeStreamResponseBaseObject, + OpenAIRealtimeStreamResponseOutputItem, OpenAIRealtimeStreamResponseOutputItemAdded, OpenAIRealtimeStreamSession, OpenAIRealtimeStreamSessionEvents, @@ -44,6 +48,17 @@ 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) @@ -1133,50 +1148,103 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): event: dict, current_output_item_id: str | None, current_response_id: str | None, - ) -> tuple[list[OpenAIRealtimeEvents], str, str, str, str]: + 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, emit the item added/delta/done trio the + OpenAI realtime protocol requires around ``function_call_arguments.done``, then close + the response so no in-progress response is left orphaned and downstream spend logging + can harvest the call from ``response.done`` output. Returns: - Tuple of (events, tool_call_id, tool_name, output_item_id, response_id) - so the caller can persist any minted IDs into session state + 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"]) response_id: Final = current_response_id or f"resp_{uuid.uuid4()}" item_id: Final = current_output_item_id or f"item_{uuid.uuid4()}" - raw_input: Final = tool_use["content"] if "content" in tool_use else tool_use.get("input") - tool_input: Final = _parse_bedrock_tool_use_input(raw_input) + tool_call_id: Final = tool_use.toolUseId + tool_name: Final = tool_use.toolName + arguments: Final = tool_use.arguments() - tool_call_id: Final = tool_use.get("toolUseId", "") - tool_name: Final = tool_use.get("toolName", "") - - 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": response_id, - "item_id": 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, - item_id, - response_id, + 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": ""} + ) + + events: Final[list[OpenAIRealtimeEvents]] = [ + 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=OpenAIRealtimeStreamResponseOutputItem({**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=OpenAIRealtimeStreamResponseOutputItem({**function_call_item}), + ), + OpenAIRealtimeDoneEvent( + type="response.done", + event_id=f"event_{uuid.uuid4()}", + response=OpenAIRealtimeResponseDoneObject( + object="realtime.response", + id=response_id, + status="completed", + output=[OpenAIRealtimeStreamResponseOutputItem({**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]: """ @@ -1381,20 +1449,22 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): ) = 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 content_end.get("type") == "TOOL": - current_response_id = None # rebind-ok: session state machine elif "toolUse" in event: - ( - events, - tool_call_id, - tool_name, - tool_output_item_id, - tool_response_id, - ) = self.transform_tool_use_event(event, current_output_item_id, current_response_id) + 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, + current_conversation_id, + ) returned_messages.extend(events) - current_output_item_id = tool_output_item_id # rebind-ok: session state machine - current_response_id = tool_response_id # rebind-ok: session state machine + # 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) diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index da0592e6bb2..c5200f0ac72 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2011,6 +2011,16 @@ class OpenAIRealtimeContentPartDone(TypedDict): type: Literal["response.content_part.done"] +class OpenAIRealtimeFunctionCallArgumentsDelta(TypedDict): + type: Literal["response.function_call_arguments.delta"] + event_id: str + response_id: str + item_id: str + output_index: int + call_id: str + delta: str + + class OpenAIRealtimeFunctionCallArgumentsDone(TypedDict): type: Literal["response.function_call_arguments.done"] event_id: str @@ -2087,6 +2097,7 @@ OpenAIRealtimeEvents = ( | OpenAIRealtimeResponseAudioDone | OpenAIRealtimeContentPartDone | OpenAIRealtimeOutputItemDone + | OpenAIRealtimeFunctionCallArgumentsDelta | OpenAIRealtimeFunctionCallArgumentsDone | OpenAIRealtimeDoneEvent ) 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 e1d5ef72c00..c91c5a35360 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 @@ -9,6 +9,7 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path +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 @@ -118,6 +119,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 @@ -364,6 +396,53 @@ class TestBedrockRealtimeSessionLifecycle: dispatched = logging_obj.dispatched_results[0] assert any(event.get("type") == "session.created" for event in dispatched) + @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() @@ -373,9 +452,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"] @@ -387,9 +464,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 6ec3c72f4bf..17bf5ae46cd 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 @@ -16,6 +16,30 @@ from litellm.llms.bedrock.realtime.transformation import ( ) from litellm.llms.bedrock.realtime.trigger_audio import ready_trigger_pcm +# The OpenAI realtime function-call lifecycle a single Nova Sonic toolUse must expand into. +_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", +] + + +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): + if event["type"] == "response.done": + return event["response"]["id"] + if event["type"] == "conversation.item.added": + return None + return event["response_id"] + class TestBedrockRealtimeConfig: """Test suite for BedrockRealtimeConfig class""" @@ -559,16 +583,11 @@ 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" - - # Verify arguments are properly formatted - args = json.loads(function_call["arguments"]) - assert args["location"] == "San Francisco" + assert json.loads(function_call["arguments"]) == {"location": "San Francisco"} def test_transform_tool_use_response_with_content_field(self): """Test toolUse response transformation with Nova 2 Sonic `content` field""" @@ -601,23 +620,34 @@ 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" + 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): - """Test transform_tool_use_event directly for input parsing and missing IDs""" + """transform_tool_use_event emits the full OpenAI function-call lifecycle""" config = BedrockRealtimeConfig() - # Missing IDs still emit a function call (Nova Sonic starts tools with role=TOOL) - events, tool_call_id, tool_name, item_id, response_id = config.transform_tool_use_event( + # 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", @@ -627,21 +657,41 @@ class TestBedrockRealtimeResponseTransformation: }, None, None, + "conv_1", ) - assert len(events) == 1 - assert events[0]["type"] == "response.function_call_arguments.done" - assert events[0]["call_id"] == "tool_call_no_ids" - assert events[0]["name"] == "get_weather" - assert response_id.startswith("resp_") - assert item_id.startswith("item_") - assert events[0]["response_id"] == response_id - assert events[0]["item_id"] == item_id - assert json.loads(events[0]["arguments"]) == {"location": "Seattle"} + assert [event["type"] for event in events] == _TOOL_CALL_EVENT_SEQUENCE + function_call = _only(events, "response.function_call_arguments.done") + 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" - # JSON string content is parsed and converted to a function call event - events, tool_call_id, tool_name, item_id, response_id = config.transform_tool_use_event( + # 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", @@ -651,19 +701,30 @@ class TestBedrockRealtimeResponseTransformation: }, "item_123", "resp_123", + "conv_1", ) - assert len(events) == 1 - assert events[0]["type"] == "response.function_call_arguments.done" - assert events[0]["call_id"] == "tool_call_123" - assert events[0]["name"] == "get_weather" - assert response_id == "resp_123" - assert item_id == "item_123" - assert events[0]["response_id"] == "resp_123" - assert events[0]["item_id"] == "item_123" - assert json.loads(events[0]["arguments"]) == {"location": "San Francisco"} + 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( + events, _, _ = config.transform_tool_use_event( { "toolUse": { "toolUseId": "tool_call_124", @@ -673,9 +734,9 @@ class TestBedrockRealtimeResponseTransformation: }, "item_123", "resp_123", + "conv_1", ) - assert len(events) == 1 - assert json.loads(events[0]["arguments"]) == {} + 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""" @@ -729,17 +790,19 @@ class TestBedrockRealtimeResponseTransformation: realtime_response_transform_input=state, ) - assert len(result["response"]) == 1 - function_call = result["response"][0] - assert function_call["type"] == "response.function_call_arguments.done" - assert result["current_response_id"] is not None - assert result["current_output_item_id"] is not None - assert result["current_response_id"].startswith("resp_") - assert result["current_output_item_id"].startswith("item_") - assert function_call["response_id"] == result["current_response_id"] - assert function_call["item_id"] == result["current_output_item_id"] + 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["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": { @@ -765,8 +828,9 @@ class TestBedrockRealtimeResponseTransformation: 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"] == [] - assert all(msg["type"] != "response.output_item.done" for msg in follow_up["response"]) post_tool_state = { "session_configuration_request": follow_up["session_configuration_request"], @@ -783,12 +847,10 @@ class TestBedrockRealtimeResponseTransformation: logging_obj, realtime_response_transform_input=post_tool_state, ) - tool_response_id = result["current_response_id"] - tool_item_id = result["current_output_item_id"] 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"] != tool_response_id - assert assistant_start["current_output_item_id"] != tool_item_id + 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"] @@ -797,7 +859,10 @@ class TestBedrockRealtimeResponseTransformation: assert added["item"]["id"] != function_call["item_id"] def test_tool_content_end_does_not_emit_message_output_item_done(self): - """Minted tool ids must not unlock unpaired message output_item.done on TOOL contentEnd""" + """ + 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" @@ -816,8 +881,8 @@ class TestBedrockRealtimeResponseTransformation: logging_obj, realtime_response_transform_input={ "session_configuration_request": json.dumps({"configured": True}), - "current_output_item_id": "item_minted_for_tool", - "current_response_id": "resp_minted_for_tool", + "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": [], @@ -826,9 +891,11 @@ class TestBedrockRealtimeResponseTransformation: ) assert result["response"] == [] - assert result["current_response_id"] is None 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""" @@ -1173,20 +1240,26 @@ class TestBedrockRealtimeContentBlockLifecycle: } }, ) - assert tool_result["response"][0]["type"] == "response.function_call_arguments.done" + 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 - self._apply( + 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 - assert state["current_delta_chunks"] is None - assert state["current_delta_type"] is None post_tool = self._apply( config, @@ -1217,6 +1290,43 @@ class TestBedrockRealtimeContentBlockLifecycle: 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_second_assistant_content_block_reuses_response_not_item(self): config = BedrockRealtimeConfig() logging_obj = MagicMock() @@ -1395,7 +1505,7 @@ class TestBedrockRealtimeUsageAccounting: assert usage["total_tokens"] == 35 assert not config.has_unbilled_usage() - def test_tool_end_turn_emits_response_done_before_clearing_ids(self): + 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" @@ -1429,6 +1539,72 @@ class TestBedrockRealtimeUsageAccounting: 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_flush_pending_usage_on_session_close(self): config = BedrockRealtimeConfig() config.record_usage_event(