mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(bedrock/realtime): close the response a Nova Sonic tool turn opens
A toolUse emitted only response.function_call_arguments.done, so the response the assistant block had already opened with response.created was never closed: the tool boundary dropped its id without a response.done and the post-tool turn minted a new one, leaving realtime clients tracking a response that stays in_progress forever. Expand a toolUse into the OpenAI function-call lifecycle the protocol expects, the same shape the Gemini realtime config emits: output_item.added, conversation.item.added, function_call_arguments.delta, .done, output_item.done, then response.done carrying the function_call item. Ids are cleared afterwards, so the tool call cannot leak into the next assistant turn and TOOL contentEnd no longer needs to drop a live response id. Because response.done now carries the call, RealTimeStreaming's shared _collect_tool_calls_from_response_done picks it up and the Bedrock-specific collector goes away. backend_ws becomes optional on RealTimeStreaming so Bedrock, which streams over the AWS SDK, no longer passes a cast bare object for it.
This commit is contained in:
parent
60822cd137
commit
06e5e7d0da
6 changed files with 464 additions and 138 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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", {})
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue