diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 858d10df53b..4ab109e5830 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -20,6 +20,8 @@ from .litellm_logging import Logging as LiteLLMLogging if TYPE_CHECKING: from websockets.asyncio.client import ClientConnection + from litellm.proxy._types import UserAPIKeyAuth + CLIENT_CONNECTION_CLASS = ClientConnection else: CLIENT_CONNECTION_CLASS = Any @@ -48,7 +50,7 @@ class RealTimeStreaming: logging_obj: LiteLLMLogging, provider_config: BaseRealtimeConfig | None = None, model: str = "", - user_api_key_dict: Any | None = None, + user_api_key_dict: "UserAPIKeyAuth | None" = None, request_data: dict | None = None, backend_uses_beta_protocol: bool | None = None, force_transcription_model: str | None = None, @@ -83,7 +85,7 @@ class RealTimeStreaming: self.current_item_chunks: list[OpenAIRealtimeOutputItemDone] | None = None self.current_delta_type: ALL_DELTA_TYPES | None = None self.session_configuration_request: str | None = None - self.user_api_key_dict = user_api_key_dict + self.user_api_key_dict: "UserAPIKeyAuth | None" = user_api_key_dict self.request_data: dict = request_data or {} # Violation counter for end_session_after_n_fails support self._violation_count: int = 0 diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index c8035a5265d..3ab0306fff5 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -9,7 +9,7 @@ store_message for backend events, store_input for client events, log_messages on import asyncio import contextlib import json -from typing import Any, Final, cast +from typing import TYPE_CHECKING, Any, Final, cast from pydantic import TypeAdapter @@ -21,6 +21,9 @@ from ..base_aws_llm import BaseAWSLLM from ..common_utils import BedrockError from .transformation import BedrockRealtimeConfig +if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth + _CLIENT_MODALITIES_ADAPTER: Final[TypeAdapter["list[str] | None"]] = TypeAdapter(list[str] | None) @@ -49,7 +52,7 @@ 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, + user_api_key_dict: "UserAPIKeyAuth | None" = None, litellm_metadata: dict | None = None, **kwargs, ): @@ -201,6 +204,11 @@ class BedrockRealtime(BaseAWSLLM): return_exceptions=True, ) finally: + for pending_usage_event in transformation_config.flush_pending_usage_as_response_done( + session_state.get("current_response_id"), + session_state.get("current_conversation_id"), + ): + realtime_streaming.store_message(pending_usage_event) await realtime_streaming.log_messages() except Exception as e: diff --git a/litellm/llms/bedrock/realtime/transformation.py b/litellm/llms/bedrock/realtime/transformation.py index 8286058bbc6..ec1e541a82e 100644 --- a/litellm/llms/bedrock/realtime/transformation.py +++ b/litellm/llms/bedrock/realtime/transformation.py @@ -36,6 +36,8 @@ from litellm.types.realtime import ( RealtimeResponseTransformInput, RealtimeResponseTypedDict, ) + + class BedrockContentEnd(BaseModel): stopReason: str | None = None @@ -168,11 +170,29 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): def record_usage_event(self, usage_event: dict) -> None: self._usage_totals = _usage_snapshot_from_event(usage_event) + def has_unbilled_usage(self) -> bool: + delta: Final = _usage_snapshot_delta(self._usage_totals, self._usage_at_last_response_done) + return any(value > 0 for value in delta.values()) + def consume_usage_for_response_done(self) -> dict[str, 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 flush_pending_usage_as_response_done( + self, + current_response_id: str | None = None, + current_conversation_id: str | None = None, + ) -> list[OpenAIRealtimeEvents]: + if not self.has_unbilled_usage(): + return [] + events, _, _, _ = self._response_done_events( + current_response_id, + current_conversation_id, + mint_ids_if_missing=True, + ) + return events + def validate_environment(self, headers: dict, model: str, api_key: str | None = None) -> dict: """Validate environment - no special validation needed for Bedrock.""" return headers @@ -830,9 +850,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): output_index=0, event_id=f"event_{uuid.uuid4()}", item_id=current_output_item_id, - part=( - {"type": "text", "text": ""} if next_delta_type == "text" else {"type": "audio", "transcript": ""} - ), + part=({"type": "text", "text": ""} if next_delta_type == "text" else {"type": "audio", "transcript": ""}), response_id=current_response_id, ) returned_messages.append(content_part_added) @@ -1051,19 +1069,27 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): Tuple of (events, reset_output_item_id, reset_response_id, reset_delta_type) """ verbose_logger.debug("Handling promptEnd") - return self._response_done_events(current_response_id, current_conversation_id) + return self._response_done_events( + current_response_id, + current_conversation_id, + mint_ids_if_missing=self.has_unbilled_usage(), + ) def _response_done_events( self, current_response_id: str | None, current_conversation_id: str | None, + *, + mint_ids_if_missing: bool = False, ) -> tuple[ list[OpenAIRealtimeEvents], str | None, str | None, ALL_DELTA_TYPES | None, ]: - if not current_response_id or not current_conversation_id: + response_id: Final = current_response_id or (f"resp_{uuid.uuid4()}" if mint_ids_if_missing else None) + conversation_id: Final = current_conversation_id or (f"conv_{uuid.uuid4()}" if mint_ids_if_missing else None) + if not response_id or not conversation_id: return [], None, None, None response_done: Final = OpenAIRealtimeDoneEvent( @@ -1071,10 +1097,10 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): event_id=f"event_{uuid.uuid4()}", response=OpenAIRealtimeResponseDoneObject( object="realtime.response", - id=current_response_id, + id=response_id, status="completed", output=[], - conversation_id=current_conversation_id, + conversation_id=conversation_id, usage=self.consume_usage_for_response_done(), ), ) @@ -1252,6 +1278,19 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): usage_event: Final = event["usageEvent"] if isinstance(usage_event, dict): self.record_usage_event(usage_event) + if current_response_id is None and self.has_unbilled_usage(): + ( + done_events, + current_output_item_id, + current_response_id, + current_delta_type, + ) = self._response_done_events( + None, + current_conversation_id, + mint_ids_if_missing=True, + ) + returned_messages.extend(done_events) + current_delta_chunks = None elif "contentStart" in event: ( @@ -1297,9 +1336,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): current_delta_chunks = None current_delta_type = None current_output_item_id = None - if content_end.get("type") == "TOOL": - current_response_id = None - if BedrockContentEnd.model_validate(content_end).stopReason == "END_TURN": + is_end_turn: Final = BedrockContentEnd.model_validate(content_end).stopReason == "END_TURN" + if is_end_turn: ( done_events, current_output_item_id, @@ -1308,6 +1346,8 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): ) = self._response_done_events(current_response_id, current_conversation_id) returned_messages.extend(done_events) current_delta_chunks = None + elif content_end.get("type") == "TOOL": + current_response_id = None elif "toolUse" in event: ( 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 2f02ef94e20..6ec3c72f4bf 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 @@ -15,7 +15,6 @@ from litellm.llms.bedrock.realtime.transformation import ( BedrockRealtimeConfig, ) from litellm.llms.bedrock.realtime.trigger_audio import ready_trigger_pcm -from litellm.types.llms.openai import OpenAIRealtimeEventTypes class TestBedrockRealtimeConfig: @@ -1356,6 +1355,125 @@ class TestBedrockRealtimeUsageAccounting: assert second_usage["input_token_details"]["audio_tokens"] == 5 assert second_usage["output_token_details"]["audio_tokens"] == 10 + def test_late_usage_event_after_response_id_cleared_emits_response_done(self): + config = BedrockRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_late_usage" + state = { + "session_configuration_request": json.dumps({"configured": True}), + "current_output_item_id": "item_1", + "current_response_id": "resp_1", + "current_conversation_id": "conv_1", + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": "audio", + } + + end_turn = config.transform_realtime_response( + json.dumps({"event": {"contentEnd": {"stopReason": "END_TURN", "type": "AUDIO"}}}), + "amazon.nova-sonic-v1:0", + logging_obj, + realtime_response_transform_input=state, + ) + assert any(msg["type"] == "response.done" for msg in end_turn["response"]) + assert end_turn["current_response_id"] is None + state["current_response_id"] = end_turn["current_response_id"] + state["current_output_item_id"] = end_turn["current_output_item_id"] + state["current_conversation_id"] = end_turn["current_conversation_id"] + + late_usage = config.transform_realtime_response( + json.dumps(self._usage_event(input_speech=10, input_text=2, output_speech=20, output_text=3)), + "amazon.nova-sonic-v1:0", + logging_obj, + realtime_response_transform_input=state, + ) + done_events = [msg for msg in late_usage["response"] if msg["type"] == "response.done"] + assert len(done_events) == 1 + usage = done_events[0]["response"]["usage"] + assert usage["input_tokens"] == 12 + assert usage["output_tokens"] == 23 + assert usage["total_tokens"] == 35 + assert not config.has_unbilled_usage() + + def test_tool_end_turn_emits_response_done_before_clearing_ids(self): + config = BedrockRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_tool_end_turn" + state = { + "session_configuration_request": json.dumps({"configured": True}), + "current_output_item_id": "item_tool", + "current_response_id": "resp_tool", + "current_conversation_id": "conv_tool", + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + } + + config.transform_realtime_response( + json.dumps(self._usage_event(input_speech=4, input_text=1, output_speech=0, output_text=2)), + "amazon.nova-sonic-v1:0", + logging_obj, + realtime_response_transform_input=state, + ) + tool_end = config.transform_realtime_response( + json.dumps({"event": {"contentEnd": {"stopReason": "END_TURN", "type": "TOOL"}}}), + "amazon.nova-sonic-v1:0", + logging_obj, + realtime_response_transform_input=state, + ) + done_events = [msg for msg in tool_end["response"] if msg["type"] == "response.done"] + assert len(done_events) == 1 + assert done_events[0]["response"]["id"] == "resp_tool" + assert done_events[0]["response"]["usage"]["input_tokens"] == 5 + assert done_events[0]["response"]["usage"]["output_tokens"] == 2 + assert tool_end["current_response_id"] is None + assert not config.has_unbilled_usage() + + def test_flush_pending_usage_on_session_close(self): + config = BedrockRealtimeConfig() + config.record_usage_event( + self._usage_event(input_speech=8, input_text=1, output_speech=16, output_text=2)["event"]["usageEvent"] + ) + assert config.has_unbilled_usage() + + flushed = config.flush_pending_usage_as_response_done(None, None) + assert len(flushed) == 1 + assert flushed[0]["type"] == "response.done" + usage = flushed[0]["response"]["usage"] + assert usage["input_tokens"] == 9 + assert usage["output_tokens"] == 18 + assert usage["total_tokens"] == 27 + assert not config.has_unbilled_usage() + assert config.flush_pending_usage_as_response_done(None, None) == [] + + def test_completion_end_with_unbilled_usage_mints_response_done(self): + config = BedrockRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_completion_end" + config.record_usage_event( + self._usage_event(input_speech=3, input_text=0, output_speech=6, output_text=0)["event"]["usageEvent"] + ) + + result = config.transform_realtime_response( + json.dumps({"event": {"completionEnd": {}}}), + "amazon.nova-sonic-v1:0", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": json.dumps({"configured": True}), + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": None, + "current_item_chunks": None, + "current_delta_type": None, + }, + ) + done_events = [msg for msg in result["response"] if msg["type"] == "response.done"] + assert len(done_events) == 1 + assert done_events[0]["response"]["usage"]["input_tokens"] == 3 + assert done_events[0]["response"]["usage"]["output_tokens"] == 6 + assert not config.has_unbilled_usage() + if __name__ == "__main__": pytest.main([__file__, "-v"])