fix(bedrock/realtime): meter usage and log via RealTimeStreaming

Parse Nova Sonic usageEvent into turn-level response.done usage, reuse
RealTimeStreaming store_message/store_input/log_messages for spend and
budget accounting, and pass user key metadata through arealtime
This commit is contained in:
mubashir1osmani 2026-08-07 15:56:37 -07:00
parent 7444897f13
commit 7d60bf6126
5 changed files with 337 additions and 41 deletions

View file

@ -2,17 +2,20 @@
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 typing import Any, Final, cast
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
@ -46,6 +49,8 @@ class BedrockRealtime(BaseAWSLLM):
aws_sts_endpoint: str | None = None,
aws_bedrock_runtime_endpoint: str | None = None,
aws_external_id: str | None = None,
user_api_key_dict: Any | None = None,
litellm_metadata: dict | None = None,
**kwargs,
):
"""
@ -120,6 +125,27 @@ class BedrockRealtime(BaseAWSLLM):
transformation_config: Final = BedrockRealtimeConfig()
logging_obj.pre_call(
input=None,
api_key=api_key or "",
additional_args={
"api_base": endpoint_uri,
"complete_input_dict": {"model": model},
},
)
# 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.
realtime_streaming: Final = RealTimeStreaming(
websocket=websocket,
backend_ws=cast(Any, object()),
logging_obj=logging_obj,
model=model,
user_api_key_dict=user_api_key_dict,
request_data={"litellm_metadata": litellm_metadata or {}},
)
try:
# Initialize the bidirectional stream
bedrock_stream: Final = await bedrock_client.invoke_model_with_bidirectional_stream(
@ -128,7 +154,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 +170,38 @@ 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:
await realtime_streaming.log_messages()
except Exception as e:
verbose_proxy_logger.exception("Error in BedrockRealtime.async_realtime: %s", e)
@ -180,6 +211,24 @@ 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(
{
"id": message.get("call_id", ""),
"type": "function",
"function": {
"name": message.get("name", ""),
"arguments": message.get("arguments", "{}"),
},
}
)
async def _forward_client_to_bedrock(
self,
client_ws: Any,
@ -188,6 +237,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 +258,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 +283,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: Final = transformation_config.session_updated_event(
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 +306,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 +359,9 @@ 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)
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])

View file

@ -36,9 +36,6 @@ from litellm.types.realtime import (
RealtimeResponseTransformInput,
RealtimeResponseTypedDict,
)
from litellm.utils import get_empty_usage
class BedrockContentEnd(BaseModel):
stopReason: str | None = None
@ -61,6 +58,74 @@ def _parse_bedrock_tool_use_input(raw_input: object) -> object:
return {}
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 _empty_usage_snapshot() -> dict[str, int]:
return {
"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: dict) -> dict[str, int]:
details: Final = usage_event.get("details") if isinstance(usage_event.get("details"), dict) else {}
total_block: Final = details.get("total") if isinstance(details.get("total"), dict) else {}
input_block: Final = total_block.get("input") if isinstance(total_block.get("input"), dict) else {}
output_block: Final = total_block.get("output") if isinstance(total_block.get("output"), dict) else {}
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 {
"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: dict[str, int], previous: dict[str, int]) -> dict[str, int]:
return {key: max(0, current.get(key, 0) - previous.get(key, 0)) for key in _empty_usage_snapshot()}
def _openai_usage_from_snapshot(snapshot: dict[str, int]) -> dict[str, Any]:
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 {
"total_tokens": snapshot.get("total", 0) or (input_tokens + output_tokens),
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"input_token_details": {
"text_tokens": snapshot.get("input_text", 0),
"audio_tokens": snapshot.get("input_speech", 0),
"cached_tokens": 0,
},
"output_token_details": {
"text_tokens": snapshot.get("output_text", 0),
"audio_tokens": snapshot.get("output_speech", 0),
},
}
class BedrockRealtimeConfig(BaseRealtimeConfig):
"""Configuration for Bedrock Nova Sonic realtime transformations."""
@ -71,6 +136,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
@ -98,6 +165,14 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
# Text configuration
self.text_media_type = "text/plain"
def record_usage_event(self, usage_event: dict) -> None:
self._usage_totals = _usage_snapshot_from_event(usage_event)
def consume_usage_for_response_done(self) -> dict[str, Any]:
delta: Final = _usage_snapshot_delta(self._usage_totals, self._usage_at_last_response_done)
self._usage_at_last_response_done = dict(self._usage_totals)
return _openai_usage_from_snapshot(delta)
def validate_environment(self, headers: dict, model: str, api_key: str | None = None) -> dict:
"""Validate environment - no special validation needed for Bedrock."""
return headers
@ -999,7 +1074,6 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
if not current_response_id or not current_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()}",
@ -1009,15 +1083,10 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
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,
},
usage=self.consume_usage_for_response_done(),
),
)
# Reset state for next response
return [response_done], None, None, None
def transform_tool_use_event(
@ -1187,6 +1256,11 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
if "sessionStart" in event:
session_configuration_request = json.dumps({"configured": True})
elif "usageEvent" in event:
usage_event: Final = event["usageEvent"]
if isinstance(usage_event, dict):
self.record_usage_event(usage_event)
elif "contentStart" in event:
(
events,

View file

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

View file

@ -57,6 +57,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:
@ -313,6 +340,30 @@ 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_session_update_is_acked_with_session_updated(self, stub_aws_models):
handler = BedrockRealtime()

View file

@ -1090,5 +1090,116 @@ class TestBedrockRealtimeSessionEvents:
assert event["session"]["modalities"] == ["text", "audio"]
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
if __name__ == "__main__":
pytest.main([__file__, "-v"])