diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index 3eeb3cb9fc6..c8035a5265d 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -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]) diff --git a/litellm/llms/bedrock/realtime/transformation.py b/litellm/llms/bedrock/realtime/transformation.py index 4d9158853cd..98716ae89c7 100644 --- a/litellm/llms/bedrock/realtime/transformation.py +++ b/litellm/llms/bedrock/realtime/transformation.py @@ -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, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index d5195659b1c..203537d689a 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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 = ( 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 ffe21b91ab2..e1d5ef72c00 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 @@ -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() 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 798bc86023d..79b662fa841 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 @@ -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"])