From 5c5699b65df42cceaaf2b5fd8782bc8675a659fa Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 17 May 2025 07:36:56 -0700 Subject: [PATCH] build: merge in https://github.com/BerriAI/litellm/pull/10909 Closes https://github.com/BerriAI/litellm/pull/10909 --- .../litellm_core_utils/realtime_streaming.py | 51 +- .../llms/base_llm/realtime/transformation.py | 9 +- litellm/llms/custom_httpx/llm_http_handler.py | 2 +- .../llms/gemini/realtime/transformation.py | 639 +++++++++++++----- .../vertex_and_google_ai_studio_gemini.py | 21 +- .../proxy/_experimental/out/onboarding.html | 1 - litellm/proxy/_new_secret_config.yaml | 4 + litellm/types/llms/gemini.py | 80 ++- litellm/types/llms/openai.py | 39 +- litellm/types/llms/vertex_ai.py | 7 +- litellm/types/realtime.py | 22 +- litellm/utils.py | 8 + .../test_gemini_realtime_transformation.py | 124 +++- ...test_vertex_and_google_ai_studio_gemini.py | 45 +- .../test_bedrock_completion.py | 2 + 15 files changed, 830 insertions(+), 224 deletions(-) delete mode 100644 litellm/proxy/_experimental/out/onboarding.html diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 347eef70a9..329f2b63c2 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -9,10 +9,11 @@ from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.types.llms.openai import ( OpenAIRealtimeEvents, OpenAIRealtimeOutputItemDone, - OpenAIRealtimeResponseTextDelta, + OpenAIRealtimeResponseDelta, OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents, ) +from litellm.types.realtime import ALL_DELTA_TYPES from .litellm_logging import Logging as LiteLLMLogging @@ -55,13 +56,13 @@ class RealTimeStreaming: self.logged_real_time_event_types = _logged_real_time_event_types self.provider_config = provider_config self.model = model - self.current_delta_chunks: Optional[ - List[OpenAIRealtimeResponseTextDelta] - ] = None + self.current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]] = None self.current_output_item_id: Optional[str] = None self.current_response_id: Optional[str] = None self.current_conversation_id: Optional[str] = None self.current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]] = None + self.current_delta_type: Optional[ALL_DELTA_TYPES] = None + self.session_configuration_request: Optional[str] = None def _should_store_message( self, @@ -112,9 +113,7 @@ class RealTimeStreaming: ## SYNC LOGGING executor.submit(self.logging_obj.success_handler(self.messages)) - async def backend_to_client_send_messages( - self, session_configuration_request: Optional[str] = None - ): + async def backend_to_client_send_messages(self): import websockets try: @@ -132,12 +131,13 @@ class RealTimeStreaming: self.model, self.logging_obj, realtime_response_transform_input={ - "session_configuration_request": session_configuration_request, + "session_configuration_request": self.session_configuration_request, "current_output_item_id": self.current_output_item_id, "current_response_id": self.current_response_id, "current_delta_chunks": self.current_delta_chunks, "current_conversation_id": self.current_conversation_id, "current_item_chunks": self.current_item_chunks, + "current_delta_type": self.current_delta_type, }, ) @@ -151,6 +151,10 @@ class RealTimeStreaming: "current_conversation_id" ] self.current_item_chunks = returned_object["current_item_chunks"] + self.current_delta_type = returned_object["current_delta_type"] + self.session_configuration_request = returned_object[ + "session_configuration_request" + ] if isinstance(transformed_response, list): for event in transformed_response: event_str = json.dumps(event) @@ -186,33 +190,20 @@ class RealTimeStreaming: self.store_input(message=message) ## FORWARD TO BACKEND if self.provider_config: - message = self.provider_config.transform_realtime_request(message) + message = self.provider_config.transform_realtime_request( + message, self.model + ) + + for msg in message: + await self.backend_ws.send(msg) + else: + await self.backend_ws.send(message) - await self.backend_ws.send(message) - except self.websocket.exceptions.ConnectionClosed: # type: ignore - verbose_logger.debug("Connection closed") - pass except Exception as e: verbose_logger.debug(f"Error in client ack messages: {e}") async def bidirectional_forward(self): - session_configuration_request: Optional[str] = None - if ( - self.provider_config - and self.provider_config.requires_session_configuration() - ): - session_configuration_request = ( - self.provider_config.session_configuration_request(self.model) - ) - if session_configuration_request is None: - raise ValueError( - "Session configuration request is None, but requires_session_configuration is True" - ) - await self.backend_ws.send(session_configuration_request) - - forward_task = asyncio.create_task( - self.backend_to_client_send_messages(session_configuration_request) - ) + forward_task = asyncio.create_task(self.backend_to_client_send_messages()) try: await self.client_ack_messages() except self.websocket.exceptions.ConnectionClosed: # type: ignore diff --git a/litellm/llms/base_llm/realtime/transformation.py b/litellm/llms/base_llm/realtime/transformation.py index db98b7e56a..d5531a532b 100644 --- a/litellm/llms/base_llm/realtime/transformation.py +++ b/litellm/llms/base_llm/realtime/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, List, Optional, Union import httpx @@ -51,7 +51,12 @@ class BaseRealtimeConfig(ABC): ) @abstractmethod - def transform_realtime_request(self, message: str) -> str: + def transform_realtime_request( + self, + message: str, + model: str, + session_configuration_request: Optional[str] = None, + ) -> List[str]: pass def requires_session_configuration( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index bb66b419b9..4cf89accfc 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2067,7 +2067,7 @@ class BaseLLMHTTPHandler: try: async with websockets.connect( # type: ignore - url, additional_headers=headers + url, extra_headers=headers ) as backend_ws: realtime_streaming = RealTimeStreaming( websocket, diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index abddc766af..59360281ac 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -7,6 +7,7 @@ import os import uuid from typing import Any, Dict, List, Optional, Union, cast +from litellm import verbose_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( @@ -16,36 +17,51 @@ from litellm.responses.litellm_completion_transformation.transformation import ( LiteLLMCompletionResponsesConfig, ) from litellm.types.llms.gemini import ( + AutomaticActivityDetection, + BidiGenerateContentRealtimeInput, + BidiGenerateContentRealtimeInputConfig, BidiGenerateContentServerContent, BidiGenerateContentServerMessage, + BidiGenerateContentSetup, ) from litellm.types.llms.openai import ( OpenAIRealtimeContentPartDone, OpenAIRealtimeConversationItemCreated, OpenAIRealtimeDoneEvent, OpenAIRealtimeEvents, + OpenAIRealtimeEventTypes, OpenAIRealtimeOutputItemDone, + OpenAIRealtimeResponseAudioDone, OpenAIRealtimeResponseContentPartAdded, + OpenAIRealtimeResponseDelta, OpenAIRealtimeResponseDoneObject, - OpenAIRealtimeResponseTextDelta, OpenAIRealtimeResponseTextDone, OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamResponseOutputItemAdded, OpenAIRealtimeStreamSession, OpenAIRealtimeStreamSessionEvents, + OpenAIRealtimeTurnDetection, +) +from litellm.types.llms.vertex_ai import ( + GeminiResponseModalities, + HttpxBlobType, + HttpxContentType, ) from litellm.types.realtime import ( + ALL_DELTA_TYPES, + RealtimeModalityResponseTransformOutput, RealtimeResponseTransformInput, RealtimeResponseTypedDict, ) +from litellm.utils import get_empty_usage from ..common_utils import encode_unserializable_types -MAP_GEMINI_FIELD_TO_OPENAI_EVENT = { - "setupComplete": "session.created", - "serverContent.modelTurn": "response.text.delta", - "serverContent.generationComplete": "response.text.done", - "serverContent.turnComplete": "response.done", +MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[str, OpenAIRealtimeEventTypes] = { + "setupComplete": OpenAIRealtimeEventTypes.SESSION_CREATED, + "serverContent.generationComplete": OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE, + "serverContent.turnComplete": OpenAIRealtimeEventTypes.RESPONSE_DONE, + "serverContent.interrupted": OpenAIRealtimeEventTypes.RESPONSE_DONE, } @@ -72,9 +88,162 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): api_base = api_base.replace("http://", "ws://") return f"{api_base}/ws/google.ai.generativelanguage.v1beta.GenerativeService.BidiGenerateContent?key={api_key}" - def transform_realtime_request(self, message: str) -> str: - realtime_input_dict: Dict[str, Any] = {} - realtime_input_dict["text"] = message + def map_model_turn_event( + self, model_turn: HttpxContentType + ) -> OpenAIRealtimeEventTypes: + """ + Map the model turn event to the OpenAI realtime events. + + Returns either: + - response.text.delta - model_turn: {"parts": [{"text": "..."}]} + - response.audio.delta - model_turn: {"parts": [{"inlineData": {"mimeType": "audio/pcm", "data": "..."}}]} + + Assumes parts is a single element list. + """ + if "parts" in model_turn: + parts = model_turn["parts"] + if len(parts) != 1: + verbose_logger.warning( + f"Realtime: Expected 1 part, got {len(parts)} for Gemini model turn event." + ) + part = parts[0] + if "text" in part: + return OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA + elif "inlineData" in part: + return OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA + else: + raise ValueError(f"Unexpected part type: {part}") + raise ValueError(f"Unexpected model turn event, no 'parts' key: {model_turn}") + + def map_generation_complete_event( + self, delta_type: Optional[ALL_DELTA_TYPES] + ) -> OpenAIRealtimeEventTypes: + if delta_type == "text": + return OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE + elif delta_type == "audio": + return OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE + else: + raise ValueError(f"Unexpected delta type: {delta_type}") + + def get_audio_mime_type(self, input_audio_format: str = "pcm16"): + mime_types = { + "pcm16": "audio/pcm", + "g711_ulaw": "audio/pcmu", + "g711_alaw": "audio/pcma", + } + + return mime_types.get(input_audio_format, "application/octet-stream") + + def map_automatic_turn_detection( + self, value: OpenAIRealtimeTurnDetection + ) -> AutomaticActivityDetection: + automatic_activity_dection = AutomaticActivityDetection() + if "create_response" in value and isinstance(value["create_response"], bool): + automatic_activity_dection["disabled"] = not value["create_response"] + else: + automatic_activity_dection["disabled"] = True + if "prefix_padding_ms" in value and isinstance(value["prefix_padding_ms"], int): + automatic_activity_dection["prefixPaddingMs"] = value["prefix_padding_ms"] + if "silence_duration_ms" in value and isinstance( + value["silence_duration_ms"], int + ): + automatic_activity_dection["silenceDurationMs"] = value[ + "silence_duration_ms" + ] + return automatic_activity_dection + + def map_openai_params( + self, optional_params: dict, non_default_params: dict + ) -> dict: + if "generationConfig" not in optional_params: + optional_params["generationConfig"] = {} + for key, value in non_default_params.items(): + if key == "instructions": + optional_params["systemInstruction"] = HttpxContentType( + role="user", parts=[{"text": value}] + ) + elif key == "temperature": + optional_params["generationConfig"]["temperature"] = value + elif key == "max_response_output_tokens": + optional_params["generationConfig"]["maxOutputTokens"] = value + elif key == "modalities": + optional_params["generationConfig"]["responseModalities"] = [ + modality.upper() for modality in cast(List[str], value) + ] + elif key == "tools": + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + vertex_gemini_config = VertexGeminiConfig() + vertex_gemini_config._map_function(value) + optional_params["generationConfig"][ + "tools" + ] = vertex_gemini_config._map_function(value) + elif key == "input_audio_transcription" and value is not None: + optional_params["inputAudioTranscription"] = {} + elif key == "turn_detection": + value_typed = cast(OpenAIRealtimeTurnDetection, value) + transformed_audio_activity_config = self.map_automatic_turn_detection( + value_typed + ) + if ( + len(transformed_audio_activity_config) > 0 + ): # if the config is not empty, add it to the optional params + optional_params[ + "realtimeInputConfig" + ] = BidiGenerateContentRealtimeInputConfig( + automaticActivityDetection=transformed_audio_activity_config + ) + if len(optional_params["generationConfig"]) == 0: + optional_params.pop("generationConfig") + return optional_params + + def transform_realtime_request( + self, + message: str, + model: str, + session_configuration_request: Optional[str] = None, + ) -> List[str]: + realtime_input_dict: BidiGenerateContentRealtimeInput = {} + try: + json_message = json.loads(message) + except json.JSONDecodeError: + if isinstance(message, bytes): + message_str = message.decode("utf-8", errors="replace") + else: + message_str = str(message) + raise ValueError(f"Invalid JSON message: {message_str}") + + ## HANDLE SESSION UPDATE ## + messages: List[str] = [] + if "type" in json_message and json_message["type"] == "session.update": + client_session_configuration_request = self.map_openai_params( + optional_params={}, non_default_params=json_message["session"] + ) + client_session_configuration_request["model"] = f"models/{model}" + + messages.append( + json.dumps( + { + "setup": client_session_configuration_request, + } + ) + ) + # elif session_configuration_request is None: + # default_session_configuration_request = self.session_configuration_request(model) + # messages.append(default_session_configuration_request) + + ## HANDLE INPUT AUDIO BUFFER ## + if ( + "type" in json_message + and json_message["type"] == "input_audio_buffer.append" + ): + realtime_input_dict["audio"] = HttpxBlobType( + mimeType=self.get_audio_mime_type(), data=json_message["audio"] + ) + else: + realtime_input_dict["text"] = message if len(realtime_input_dict) != 1: raise ValueError( @@ -82,9 +251,13 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): f" {list(realtime_input_dict.keys())}" ) - realtime_input_dict = encode_unserializable_types(realtime_input_dict) + realtime_input_dict = cast( + BidiGenerateContentRealtimeInput, + encode_unserializable_types(cast(Dict[str, object], realtime_input_dict)), + ) - return json.dumps({"realtime_input": realtime_input_dict}) + messages.append(json.dumps({"realtime_input": realtime_input_dict})) + return messages def transform_session_created_event( self, @@ -92,16 +265,21 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): logging_session_id: str, session_configuration_request: Optional[str] = None, ) -> OpenAIRealtimeStreamSessionEvents: - if session_configuration_request is None: - raise ValueError( - "session_configuration_request is required for Gemini API calls" - ) + if session_configuration_request: + session_configuration_request_dict: BidiGenerateContentSetup = json.loads( + session_configuration_request + ).get("setup", {}) + else: + session_configuration_request_dict = {} - session_configuration_request_dict = json.loads(session_configuration_request) _model = session_configuration_request_dict.get("model") or model - _modalities = session_configuration_request_dict.get( - "generationConfig", {} - ).get("responseModalities", ["TEXT"]) + generation_config = ( + session_configuration_request_dict.get("generationConfig", {}) or {} + ) + gemini_modalities = generation_config.get("responseModalities", ["TEXT"]) + _modalities = [ + modality.lower() for modality in cast(List[str], gemini_modalities) + ] _system_instruction = session_configuration_request_dict.get( "systemInstruction" ) @@ -112,7 +290,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): if _system_instruction is not None and isinstance(_system_instruction, str): session["instructions"] = _system_instruction if _model is not None and isinstance(_model, str): - session["model"] = _model + session["model"] = _model.strip( + "models/" + ) # keep it consistent with how openai returns the model name return OpenAIRealtimeStreamSessionEvents( type="session.created", @@ -137,6 +317,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): response_id: str, output_item_id: str, conversation_id: str, + delta_type: ALL_DELTA_TYPES, session_configuration_request: Optional[str] = None, ) -> List[OpenAIRealtimeEvents]: if session_configuration_request is None: @@ -144,16 +325,19 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): "session_configuration_request is required for Gemini API calls" ) - session_configuration_request_dict = json.loads(session_configuration_request) - _modalities = session_configuration_request_dict.get( + session_configuration_request_dict: BidiGenerateContentSetup = json.loads( + session_configuration_request + ).get("setup", {}) + generation_config = session_configuration_request_dict.get( "generationConfig", {} - ).get("responseModalities", ["TEXT"]) - _temperature = session_configuration_request_dict.get( - "generationConfig", {} - ).get("temperature") - _max_output_tokens = session_configuration_request_dict.get( - "generationConfig", {} - ).get("maxOutputTokens") + ) + gemini_modalities = generation_config.get("responseModalities", ["TEXT"]) + _modalities = [ + modality.lower() for modality in cast(List[str], gemini_modalities) + ] + + _temperature = generation_config.get("temperature") + _max_output_tokens = generation_config.get("maxOutputTokens") response_items: List[OpenAIRealtimeEvents] = [] @@ -213,6 +397,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): part={ "type": "text", "text": "", + } + if delta_type == "text" + else { + "type": "audio", + "transcript": "", }, response_id=response_id, ) @@ -224,20 +413,25 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): message: BidiGenerateContentServerContent, output_item_id: str, response_id: str, - ) -> OpenAIRealtimeResponseTextDelta: + delta_type: ALL_DELTA_TYPES, + ) -> OpenAIRealtimeResponseDelta: delta = "" try: if "modelTurn" in message and "parts" in message["modelTurn"]: for part in message["modelTurn"]["parts"]: if "text" in part: delta += part["text"] + elif "inlineData" in part: + delta += part["inlineData"]["data"] except Exception as e: raise ValueError( f"Error transforming content delta events: {e}, got message: {message}" ) - return OpenAIRealtimeResponseTextDelta( - type="response.text.delta", + return OpenAIRealtimeResponseDelta( + type="response.text.delta" + if delta_type == "text" + else "response.audio.delta", content_index=0, event_id="event_{}".format(uuid.uuid4()), item_id=output_item_id, @@ -248,10 +442,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def transform_content_done_event( self, - delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]], + delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]], current_output_item_id: Optional[str], current_response_id: Optional[str], - ) -> OpenAIRealtimeResponseTextDone: + delta_type: ALL_DELTA_TYPES, + ) -> Union[OpenAIRealtimeResponseTextDone, OpenAIRealtimeResponseAudioDone]: if delta_chunks: delta = "".join([delta_chunk["delta"] for delta_chunk in delta_chunks]) else: @@ -260,21 +455,34 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): raise ValueError( "current_output_item_id and current_response_id cannot be None for a 'done' event." ) - return OpenAIRealtimeResponseTextDone( - type="response.text.done", - content_index=0, - event_id="event_{}".format(uuid.uuid4()), - item_id=current_output_item_id, - output_index=0, - response_id=current_response_id, - text=delta, - ) + if delta_type == "text": + return OpenAIRealtimeResponseTextDone( + type="response.text.done", + content_index=0, + event_id="event_{}".format(uuid.uuid4()), + item_id=current_output_item_id, + output_index=0, + response_id=current_response_id, + text=delta, + ) + elif delta_type == "audio": + return OpenAIRealtimeResponseAudioDone( + type="response.audio.done", + content_index=0, + event_id="event_{}".format(uuid.uuid4()), + item_id=current_output_item_id, + output_index=0, + response_id=current_response_id, + ) def return_additional_content_done_events( self, current_output_item_id: Optional[str], current_response_id: Optional[str], - delta_done_event: OpenAIRealtimeResponseTextDone, + delta_done_event: Union[ + OpenAIRealtimeResponseTextDone, OpenAIRealtimeResponseAudioDone + ], + delta_type: ALL_DELTA_TYPES, ) -> List[OpenAIRealtimeEvents]: """ - return response.content_part.done @@ -285,6 +493,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): "current_output_item_id and current_response_id cannot be None for a 'done' event." ) returned_items: List[OpenAIRealtimeEvents] = [] + + delta_done_event_text = cast(Optional[str], delta_done_event.get("text")) # response.content_part.done response_content_part_done = OpenAIRealtimeContentPartDone( type="response.content_part.done", @@ -292,9 +502,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): event_id="event_{}".format(uuid.uuid4()), item_id=current_output_item_id, output_index=0, - part={ - "type": "text", - "text": delta_done_event["text"], + part={"type": "text", "text": delta_done_event_text} + if delta_done_event_text and delta_type == "text" + else { + "type": "audio", + "transcript": "", # gemini doesn't return transcript for audio }, response_id=current_response_id, ) @@ -312,9 +524,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): "status": "completed", "role": "assistant", "content": [ - { - "type": "text", - "text": delta_done_event["text"], + {"type": "text", "text": delta_done_event_text} + if delta_done_event_text and delta_type == "text" + else { + "type": "audio", + "transcript": "", } ], }, @@ -336,8 +550,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def update_current_delta_chunks( self, transformed_message: Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]], - current_delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]], - ) -> Optional[List[OpenAIRealtimeResponseTextDelta]]: + current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]], + ) -> Optional[List[OpenAIRealtimeResponseDelta]]: try: if isinstance(transformed_message, list): current_delta_chunks = [] @@ -345,7 +559,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): for event in transformed_message: if event["type"] == "response.text.delta": current_delta_chunks.append( - cast(OpenAIRealtimeResponseTextDelta, event) + cast(OpenAIRealtimeResponseDelta, event) ) any_delta_chunk = True if not any_delta_chunk: @@ -353,11 +567,13 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): None # reset current_delta_chunks if no delta chunks ) else: - if transformed_message["type"] == "response.text.delta": + if ( + transformed_message["type"] == "response.text.delta" + ): # ONLY ACCUMULATE TEXT DELTA CHUNKS - AUDIO WILL CAUSE SERVER MEMORY ISSUES if current_delta_chunks is None: current_delta_chunks = [] current_delta_chunks.append( - cast(OpenAIRealtimeResponseTextDelta, transformed_message) + cast(OpenAIRealtimeResponseDelta, transformed_message) ) else: current_delta_chunks = None @@ -406,40 +622,41 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): message: BidiGenerateContentServerMessage, current_response_id: Optional[str], current_conversation_id: Optional[str], - current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]], output_items: Optional[List[OpenAIRealtimeOutputItemDone]], session_configuration_request: Optional[str] = None, ) -> OpenAIRealtimeDoneEvent: - if ( - current_conversation_id is None - or current_response_id is None - or current_item_chunks is None - ): + if current_conversation_id is None or current_response_id is None: raise ValueError( - "current_conversation_id and current_response_id and current_item_chunks cannot be None for a 'done' event." - ) - if session_configuration_request is None: - raise ValueError( - "session_configuration_request is required for Gemini API calls" + f"current_conversation_id and current_response_id must all be set for a 'done' event. Got=current_conversation_id: {current_conversation_id}, current_response_id: {current_response_id}" ) - session_configuration_request_dict = json.loads(session_configuration_request) - temperature = session_configuration_request_dict.get( + if session_configuration_request: + session_configuration_request_dict: BidiGenerateContentSetup = json.loads( + session_configuration_request + ).get("setup", {}) + else: + session_configuration_request_dict = {} + + generation_config = session_configuration_request_dict.get( "generationConfig", {} - ).get("temperature") - max_output_tokens = session_configuration_request_dict.get( - "generationConfig", {} - ).get("maxOutputTokens") - _modalities = session_configuration_request_dict.get( - "generationConfig", {} - ).get("responseModalities", ["TEXT"]) - _chat_completion_usage = VertexGeminiConfig()._calculate_usage( - completion_response=message, ) + temperature = generation_config.get("temperature") + max_output_tokens = generation_config.get("max_output_tokens") + gemini_modalities = generation_config.get("responseModalities", ["TEXT"]) + _modalities = [ + modality.lower() for modality in cast(List[str], gemini_modalities) + ] + if "usageMetadata" in message: + _chat_completion_usage = VertexGeminiConfig()._calculate_usage( + completion_response=message, + ) + else: + _chat_completion_usage = get_empty_usage() + responses_api_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( _chat_completion_usage, ) - return OpenAIRealtimeDoneEvent( + response_done_event = OpenAIRealtimeDoneEvent( type="response.done", event_id="event_{}".format(uuid.uuid4()), response=OpenAIRealtimeResponseDoneObject( @@ -451,11 +668,121 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): else [], conversation_id=current_conversation_id, modalities=_modalities, - temperature=temperature, - max_output_tokens=max_output_tokens, usage=responses_api_usage.model_dump(), ), ) + if temperature is not None: + response_done_event["response"]["temperature"] = temperature + if max_output_tokens is not None: + response_done_event["response"]["max_output_tokens"] = max_output_tokens + + return response_done_event + + def handle_openai_modality_event( + self, + openai_event: OpenAIRealtimeEventTypes, + json_message: dict, + realtime_response_transform_input: RealtimeResponseTransformInput, + delta_type: ALL_DELTA_TYPES, + ) -> RealtimeModalityResponseTransformOutput: + current_output_item_id = realtime_response_transform_input[ + "current_output_item_id" + ] + current_response_id = realtime_response_transform_input["current_response_id"] + current_conversation_id = realtime_response_transform_input[ + "current_conversation_id" + ] + current_delta_chunks = realtime_response_transform_input["current_delta_chunks"] + session_configuration_request = realtime_response_transform_input[ + "session_configuration_request" + ] + + returned_message: List[OpenAIRealtimeEvents] = [] + if ( + openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA + or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA + ): + current_response_id = current_response_id or "resp_{}".format(uuid.uuid4()) + if not current_output_item_id: + # send the list of standard 'new' content.delta events + current_output_item_id = "item_{}".format(uuid.uuid4()) + current_conversation_id = current_conversation_id or "conv_{}".format( + uuid.uuid4() + ) + returned_message = self.return_new_content_delta_events( + session_configuration_request=session_configuration_request, + response_id=current_response_id, + output_item_id=current_output_item_id, + conversation_id=current_conversation_id, + delta_type=delta_type, + ) + + # send the list of standard 'new' content.delta events + transformed_message = self.transform_content_delta_events( + BidiGenerateContentServerContent(**json_message["serverContent"]), + current_output_item_id, + current_response_id, + delta_type=delta_type, + ) + returned_message.append(transformed_message) + elif ( + openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE + or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE + ): + transformed_content_done_event = self.transform_content_done_event( + current_output_item_id=current_output_item_id, + current_response_id=current_response_id, + delta_chunks=current_delta_chunks, + delta_type=delta_type, + ) + returned_message = [transformed_content_done_event] + + additional_items = self.return_additional_content_done_events( + current_output_item_id=current_output_item_id, + current_response_id=current_response_id, + delta_done_event=transformed_content_done_event, + delta_type=delta_type, + ) + returned_message.extend(additional_items) + + return { + "returned_message": returned_message, + "current_output_item_id": current_output_item_id, + "current_response_id": current_response_id, + "current_conversation_id": current_conversation_id, + "current_delta_chunks": current_delta_chunks, + "current_delta_type": delta_type, + } + + def map_openai_event( + self, + key: str, + value: dict, + current_delta_type: Optional[ALL_DELTA_TYPES], + json_message: dict, + ) -> OpenAIRealtimeEventTypes: + model_turn_event = value.get("modelTurn") + generation_complete_event = value.get("generationComplete") + openai_event: Optional[OpenAIRealtimeEventTypes] = None + if model_turn_event: # check if model turn event + openai_event = self.map_model_turn_event(model_turn_event) + elif generation_complete_event: + openai_event = self.map_generation_complete_event( + delta_type=current_delta_type + ) + else: + # Check if this key or any nested key matches our mapping + for map_key, openai_event in MAP_GEMINI_FIELD_TO_OPENAI_EVENT.items(): + if map_key == key or ( + "." in map_key + and GeminiRealtimeConfig.get_nested_value(json_message, map_key) + is not None + ): + openai_event = openai_event + break + if openai_event is None: + raise ValueError(f"Unknown openai event: {key}, value: {value}") + return openai_event def transform_realtime_response( self, @@ -490,91 +817,58 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): "session_configuration_request" ] current_item_chunks = realtime_response_transform_input["current_item_chunks"] - returned_message: Optional[ - Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]] - ] = None + current_delta_type: Optional[ + ALL_DELTA_TYPES + ] = realtime_response_transform_input["current_delta_type"] + returned_message: List[OpenAIRealtimeEvents] = [] + for key, value in json_message.items(): # Check if this key or any nested key matches our mapping - for map_key, openai_event in MAP_GEMINI_FIELD_TO_OPENAI_EVENT.items(): - if map_key == key or ( - "." in map_key - and GeminiRealtimeConfig.get_nested_value(json_message, map_key) - is not None - ): - if openai_event == "session.created": - transformed_message = self.transform_session_created_event( - model, - logging_session_id, - realtime_response_transform_input[ - "session_configuration_request" - ], - ) - returned_message = transformed_message + openai_event = self.map_openai_event( + key=key, + value=value, + current_delta_type=current_delta_type, + json_message=json_message, + ) - elif openai_event == "response.text.delta": - # check if this is a new content.delta or a continuation of a previous content.delta - if not current_output_item_id: - # send the list of standard 'new' content.delta events - current_response_id = ( - current_response_id or "resp_{}".format(uuid.uuid4()) - ) - current_output_item_id = "item_{}".format(uuid.uuid4()) - current_conversation_id = ( - current_conversation_id - or "conv_{}".format(uuid.uuid4()) - ) - response_items = self.return_new_content_delta_events( - session_configuration_request=session_configuration_request, - response_id=current_response_id, - output_item_id=current_output_item_id, - conversation_id=current_conversation_id, - ) - - transformed_message = self.transform_content_delta_events( - BidiGenerateContentServerContent(**json_message[key]), # type: ignore - current_output_item_id, - current_response_id, - ) - response_items.append(transformed_message) - returned_message = response_items - else: - current_response_id = ( - current_response_id or "resp_{}".format(uuid.uuid4()) - ) - # send the list of standard 'new' content.delta events - transformed_message = self.transform_content_delta_events( - BidiGenerateContentServerContent(**json_message[key]), # type: ignore - current_output_item_id, - current_response_id, - ) - returned_message = transformed_message - elif openai_event == "response.text.done": - transformed_content_done_event = ( - self.transform_content_done_event( - current_output_item_id=current_output_item_id, - current_response_id=current_response_id, - delta_chunks=current_delta_chunks, - ) - ) - returned_message = [transformed_content_done_event] - - additional_items = self.return_additional_content_done_events( - current_output_item_id=current_output_item_id, - current_response_id=current_response_id, - delta_done_event=transformed_content_done_event, - ) - returned_message.extend(additional_items) - elif openai_event == "response.done": - transformed_response_done_event = self.transform_response_done_event( - message=BidiGenerateContentServerMessage(**json_message), # type: ignore - current_response_id=current_response_id, - current_conversation_id=current_conversation_id, - session_configuration_request=session_configuration_request, - output_items=current_item_chunks, - ) - returned_message = transformed_response_done_event - - if returned_message is None: + if openai_event == OpenAIRealtimeEventTypes.SESSION_CREATED: + transformed_message = self.transform_session_created_event( + model, + logging_session_id, + realtime_response_transform_input["session_configuration_request"], + ) + session_configuration_request = json.dumps(transformed_message) + returned_message.append(transformed_message) + elif openai_event == OpenAIRealtimeEventTypes.RESPONSE_DONE: + transformed_response_done_event = self.transform_response_done_event( + message=BidiGenerateContentServerMessage(**json_message), # type: ignore + current_response_id=current_response_id, + current_conversation_id=current_conversation_id, + session_configuration_request=session_configuration_request, + output_items=None, + ) + returned_message.append(transformed_response_done_event) + elif ( + openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA + or openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE + or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA + or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE + ): + _returned_message = self.handle_openai_modality_event( + openai_event, + json_message, + realtime_response_transform_input, + delta_type="text" if "text" in openai_event.value else "audio", + ) + returned_message.extend(_returned_message["returned_message"]) + current_output_item_id = _returned_message["current_output_item_id"] + current_response_id = _returned_message["current_response_id"] + current_conversation_id = _returned_message["current_conversation_id"] + current_delta_chunks = _returned_message["current_delta_chunks"] + current_delta_type = _returned_message["current_delta_type"] + else: + raise ValueError(f"Unknown openai event: {openai_event}") + if len(returned_message) == 0: if isinstance(message, bytes): message_str = message.decode("utf-8", errors="replace") else: @@ -596,12 +890,14 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): "current_delta_chunks": current_delta_chunks, "current_conversation_id": current_conversation_id, "current_item_chunks": current_item_chunks, + "current_delta_type": current_delta_type, + "session_configuration_request": session_configuration_request, } def requires_session_configuration(self) -> bool: return True - def session_configuration_request(self, model: str) -> Optional[str]: + def session_configuration_request(self, model: str) -> str: """ ``` @@ -624,11 +920,20 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): } ``` """ + + response_modalities: List[GeminiResponseModalities] = ["AUDIO"] + output_audio_transcription = False + # if "audio" in model: ## UNCOMMENT THIS WHEN AUDIO IS SUPPORTED + # output_audio_transcription = True + + setup_config: BidiGenerateContentSetup = { + "model": f"models/{model}", + "generationConfig": {"responseModalities": response_modalities}, + } + if output_audio_transcription: + setup_config["outputAudioTranscription"] = {} return json.dumps( { - "setup": { - "model": f"models/{model}", - "generationConfig": {"responseModalities": ["TEXT"]}, - } + "setup": setup_config, } ) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 203c563436..cd67be3545 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -63,6 +63,7 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import ( ChatCompletionTokenLogprob, ChoiceLogprobs, + CompletionTokensDetailsWrapper, GenericStreamingChunk, PromptTokensDetailsWrapper, TopLogprob, @@ -803,10 +804,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): text_tokens: Optional[int] = None prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None reasoning_tokens: Optional[int] = None + response_tokens: Optional[int] = None + response_tokens_details: Optional[CompletionTokensDetailsWrapper] = None if "cachedContentTokenCount" in completion_response["usageMetadata"]: cached_tokens = completion_response["usageMetadata"][ "cachedContentTokenCount" ] + + ## GEMINI LIVE API ONLY PARAMS ## + if "responseTokenCount" in completion_response["usageMetadata"]: + response_tokens = completion_response["usageMetadata"]["responseTokenCount"] + if "responseTokensDetails" in completion_response["usageMetadata"]: + response_tokens_details = CompletionTokensDetailsWrapper() + for detail in completion_response["usageMetadata"]["responseTokensDetails"]: + if detail["modality"] == "TEXT": + response_tokens_details.text_tokens = detail["tokenCount"] + elif detail["modality"] == "AUDIO": + response_tokens_details.audio_tokens = detail["tokenCount"] + ######################################################### + if "promptTokensDetails" in completion_response["usageMetadata"]: for detail in completion_response["usageMetadata"]["promptTokensDetails"]: if detail["modality"] == "AUDIO": @@ -823,7 +839,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): text_tokens=text_tokens, ) - completion_tokens = completion_response["usageMetadata"].get( + completion_tokens = response_tokens or completion_response["usageMetadata"].get( "candidatesTokenCount", 0 ) if ( @@ -842,6 +858,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): total_tokens=completion_response["usageMetadata"].get("totalTokenCount", 0), prompt_tokens_details=prompt_tokens_details, reasoning_tokens=reasoning_tokens, + completion_tokens_details=response_tokens_details, ) return usage @@ -1637,7 +1654,7 @@ class ModelResponseIterator: "reasoning_tokens": processed_chunk["usageMetadata"].get( "thoughtsTokenCount", 0 ) - } + }, ) returned_chunk = GenericStreamingChunk( diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index 94024dc83f..0000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 6ddb8c5d06..98d09d7f8d 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -6,6 +6,10 @@ model_list: litellm_params: model: gpt-4o-mini api_key: os.environ/OPENAI_API_KEY + - model_name: "gpt-4o-realtime-preview" + litellm_params: + model: gpt-4o-realtime-preview-2024-10-01 + api_key: os.environ/OPENAI_API_KEY - model_name: "bedrock-nova" litellm_params: model: us.amazon.nova-pro-v1:0 diff --git a/litellm/types/llms/gemini.py b/litellm/types/llms/gemini.py index c381440881..e39a2a8e82 100644 --- a/litellm/types/llms/gemini.py +++ b/litellm/types/llms/gemini.py @@ -3,7 +3,13 @@ from typing import Any, Dict, Iterable, List, Literal, Optional, Union from typing_extensions import Required, TypedDict -from .vertex_ai import HttpxContentType, UsageMetadata +from .vertex_ai import ( + GenerationConfig, + HttpxBlobType, + HttpxContentType, + Tools, + UsageMetadata, +) class GeminiFilesState(Enum): @@ -72,3 +78,75 @@ class BidiGenerateContentServerMessage(TypedDict, total=False): setupComplete: dict """Output only. The setup complete message.""" + + +class BidiGenerateContentRealtimeInput(TypedDict, total=False): + text: str + """The text to be sent to the model.""" + + audio: HttpxBlobType + """The audio to be sent to the model.""" + + video: HttpxBlobType + """The video to be sent to the model.""" + + audioStreamEnd: bool + """Output only. If true, indicates that the audio stream has ended.""" + + activityStart: bool + """Output only. If true, indicates that the activity has started.""" + + activityEnd: bool + """Output only. If true, indicates that the activity has ended.""" + + +StartOfSpeechSensitivityEnum = Literal[ + "START_SENSITIVITY_UNSPECIFIED", "START_SENSITIVITY_HIGH", "START_SENSITIVITY_LOW" +] +EndOfSpeechSensitivityEnum = Literal[ + "END_SENSITIVITY_UNSPECIFIED", "END_SENSITIVITY_HIGH", "END_SENSITIVITY_LOW" +] + + +class AutomaticActivityDetection(TypedDict, total=False): + disabled: bool + startOfSpeechSensitivity: StartOfSpeechSensitivityEnum + prefixPaddingMs: int + endOfSpeechSensitivity: EndOfSpeechSensitivityEnum + silenceDurationMs: int + + +class BidiGenerateContentRealtimeInputConfig(TypedDict, total=False): + automaticActivityDetection: AutomaticActivityDetection + + +class BidiGenerateContentSetup(TypedDict, total=False): + model: str + """The model to be used for the realtime session.""" + + generationConfig: GenerationConfig + """The generation config to be used for the realtime session.""" + + systemInstruction: HttpxContentType + """The system instruction to be used for the realtime session.""" + + tools: List[Tools] + """The tools to be used for the realtime session.""" + + realtimeInputConfig: dict + """The realtime config to be used for the realtime session.""" + + sessionResumption: dict + """The session resumption to be used for the realtime session.""" + + sessionResumptionConfig: dict + """The session resumption config to be used for the realtime session.""" + + contextWindowCompression: dict + """The context window compression to be used for the realtime session.""" + + inputAudioTranscription: dict + """The input audio transcription to be used for the realtime session.""" + + outputAudioTranscription: dict + """The output audio transcription to be used for the realtime session.""" diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 8ae059d706..0d880a4b1c 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1353,7 +1353,7 @@ class OpenAIRealtimeStreamResponseOutputItemContent(TypedDict, total=False): """The text content, used for 'input_text' and 'text' content types""" transcript: str """The transcript content, used for 'input_audio' content types""" - type: Literal["input_audio", "input_text", "text", "item_reference"] + type: Literal["input_audio", "input_text", "text", "item_reference", "audio"] """The type of content""" @@ -1443,14 +1443,14 @@ class OpenAIRealtimeResponseContentPartAdded(TypedDict): response_id: str -class OpenAIRealtimeResponseTextDelta(TypedDict): +class OpenAIRealtimeResponseDelta(TypedDict): content_index: int delta: str event_id: str item_id: str output_index: int response_id: str - type: Literal["response.text.delta"] + type: Union[Literal["response.text.delta"], Literal["response.audio.delta"]] class OpenAIRealtimeResponseTextDone(TypedDict): @@ -1463,6 +1463,15 @@ class OpenAIRealtimeResponseTextDone(TypedDict): type: Literal["response.text.done"] +class OpenAIRealtimeResponseAudioDone(TypedDict): + content_index: int + event_id: str + item_id: str + output_index: int + response_id: str + type: Literal["response.audio.done"] + + class OpenAIRealtimeContentPartDone(TypedDict): content_index: int event_id: str @@ -1503,6 +1512,17 @@ class OpenAIRealtimeDoneEvent(TypedDict): type: Literal["response.done"] +class OpenAIRealtimeEventTypes(Enum): + SESSION_CREATED = "session.created" + RESPONSE_TEXT_DELTA = "response.text.delta" + RESPONSE_AUDIO_DELTA = "response.audio.delta" + RESPONSE_TEXT_DONE = "response.text.done" + RESPONSE_AUDIO_DONE = "response.audio.done" + RESPONSE_DONE = "response.done" + RESPONSE_OUTPUT_ITEM_ADDED = "response.output_item.added" + RESPONSE_CONTENT_PART_ADDED = "response.content_part.added" + + OpenAIRealtimeEvents = Union[ OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents, @@ -1510,8 +1530,9 @@ OpenAIRealtimeEvents = Union[ OpenAIRealtimeResponseContentPartAdded, OpenAIRealtimeConversationItemCreated, OpenAIRealtimeConversationCreated, - OpenAIRealtimeResponseTextDelta, + OpenAIRealtimeResponseDelta, OpenAIRealtimeResponseTextDone, + OpenAIRealtimeResponseAudioDone, OpenAIRealtimeContentPartDone, OpenAIRealtimeOutputItemDone, OpenAIRealtimeDoneEvent, @@ -1610,3 +1631,13 @@ class OpenAIWebSearchUserLocation(TypedDict): class OpenAIWebSearchOptions(TypedDict, total=False): search_context_size: Optional[Literal["low", "medium", "high"]] user_location: Optional[OpenAIWebSearchUserLocation] + + +class OpenAIRealtimeTurnDetection(TypedDict, total=False): + create_response: bool + eagerness: str + interrupt_response: bool + prefix_padding_ms: int + silence_duration_ms: int + threshold: int + type: str diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 6ac7ae6275..be43a7969e 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -173,6 +173,9 @@ class GeminiThinkingConfig(TypedDict, total=False): thinkingBudget: int +GeminiResponseModalities = Literal["TEXT", "IMAGE", "AUDIO", "VIDEO"] + + class GenerationConfig(TypedDict, total=False): temperature: float top_p: float @@ -187,7 +190,7 @@ class GenerationConfig(TypedDict, total=False): seed: int responseLogprobs: bool logprobs: int - responseModalities: List[Literal["TEXT", "IMAGE", "AUDIO", "VIDEO"]] + responseModalities: List[GeminiResponseModalities] thinkingConfig: GeminiThinkingConfig @@ -218,9 +221,11 @@ class UsageMetadata(TypedDict, total=False): promptTokenCount: int totalTokenCount: int candidatesTokenCount: int + responseTokenCount: int cachedContentTokenCount: int promptTokensDetails: List[PromptTokensDetails] thoughtsTokenCount: int + responseTokensDetails: List[PromptTokensDetails] class CachedContent(TypedDict, total=False): diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index f8d613f7f0..c105983b1e 100644 --- a/litellm/types/realtime.py +++ b/litellm/types/realtime.py @@ -1,11 +1,13 @@ -from typing import List, Optional, TypedDict, Union +from typing import List, Literal, Optional, TypedDict, Union from .llms.openai import ( OpenAIRealtimeEvents, OpenAIRealtimeOutputItemDone, - OpenAIRealtimeResponseTextDelta, + OpenAIRealtimeResponseDelta, ) +ALL_DELTA_TYPES = Literal["text", "audio"] + class RealtimeResponseTransformInput(TypedDict): session_configuration_request: Optional[str] @@ -15,15 +17,27 @@ class RealtimeResponseTransformInput(TypedDict): current_response_id: Optional[ str ] # used to check if this is a new content.delta or a continuation of a previous content.delta - current_delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]] + current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]] current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]] current_conversation_id: Optional[str] + current_delta_type: Optional[ALL_DELTA_TYPES] class RealtimeResponseTypedDict(TypedDict): response: Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]] current_output_item_id: Optional[str] current_response_id: Optional[str] - current_delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]] + current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]] current_conversation_id: Optional[str] current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]] + current_delta_type: Optional[ALL_DELTA_TYPES] + session_configuration_request: Optional[str] + + +class RealtimeModalityResponseTransformOutput(TypedDict): + returned_message: List[OpenAIRealtimeEvents] + current_output_item_id: Optional[str] + current_response_id: Optional[str] + current_conversation_id: Optional[str] + current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]] + current_delta_type: Optional[ALL_DELTA_TYPES] diff --git a/litellm/utils.py b/litellm/utils.py index 3e9a9bdfdf..1647994047 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6876,3 +6876,11 @@ def jsonify_tools(tools: List[Any]) -> List[Dict]: if isinstance(tool, dict): new_tools.append(tool) return new_tools + + +def get_empty_usage() -> Usage: + return Usage( + prompt_tokens=0, + completion_tokens=0, + total_tokens=0, + ) diff --git a/tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py b/tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py index 18dd35aba2..fe4b4583d2 100644 --- a/tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py +++ b/tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py @@ -40,9 +40,12 @@ def test_gemini_realtime_transformation_session_created(): "current_conversation_id": None, "current_delta_chunks": [], "current_item_chunks": [], + "current_delta_type": None, }, ) - assert transformed_message["response"]["type"] == "session.created" + + print(transformed_message) + assert transformed_message["response"][0]["type"] == "session.created" def test_gemini_realtime_transformation_content_delta(): @@ -80,6 +83,7 @@ def test_gemini_realtime_transformation_content_delta(): "current_conversation_id": None, "current_delta_chunks": [], "current_item_chunks": [], + "current_delta_type": None, }, ) transformed_message = returned_object["response"] @@ -105,3 +109,121 @@ def test_gemini_realtime_transformation_content_delta(): event["item_id"] for event in transformed_message if "item_id" in event ] assert len(set(output_item_ids)) == 1 + + +def test_gemini_model_turn_event_mapping(): + from litellm.types.llms.openai import OpenAIRealtimeEventTypes + + config = GeminiRealtimeConfig() + assert config is not None + + model_turn_event = {"parts": [{"text": "Hello, world!"}]} + openai_event = config.map_model_turn_event(model_turn_event) + assert openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA + + model_turn_event = { + "parts": [{"inlineData": {"mimeType": "audio/pcm", "data": "..."}}] + } + openai_event = config.map_model_turn_event(model_turn_event) + assert openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA + + model_turn_event = { + "parts": [ + { + "text": "Hello, world!", + "inlineData": {"mimeType": "audio/pcm", "data": "..."}, + } + ] + } + openai_event = config.map_model_turn_event(model_turn_event) + assert openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA + + +def test_gemini_realtime_transformation_audio_delta(): + from litellm.types.llms.openai import OpenAIRealtimeEventTypes + + config = GeminiRealtimeConfig() + assert config is not None + + session_configuration_request = { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } + session_configuration_request_str = json.dumps(session_configuration_request) + + audio_delta_event = { + "serverContent": { + "modelTurn": { + "parts": [ + {"inlineData": {"mimeType": "audio/pcm", "data": "my-audio-data"}} + ] + } + } + } + + result = config.transform_realtime_response( + json.dumps(audio_delta_event), + "gemini-1.5-flash", + MagicMock(), + realtime_response_transform_input={ + "session_configuration_request": session_configuration_request_str, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + print(result) + + responses = result["response"] + + contains_audio_delta = False + for response in responses: + if response["type"] == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA.value: + contains_audio_delta = True + break + assert contains_audio_delta, "Expected audio delta event" + + +def test_gemini_realtime_transformation_generation_complete(): + from litellm.types.llms.openai import OpenAIRealtimeEventTypes + + config = GeminiRealtimeConfig() + assert config is not None + + session_configuration_request = { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } + session_configuration_request_str = json.dumps(session_configuration_request) + + audio_delta_event = {"serverContent": {"generationComplete": True}} + + result = config.transform_realtime_response( + json.dumps(audio_delta_event), + "gemini-1.5-flash", + MagicMock(), + realtime_response_transform_input={ + "session_configuration_request": session_configuration_request_str, + "current_output_item_id": "my-output-item-id", + "current_response_id": "my-response-id", + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": "audio", + }, + ) + + print(result) + + responses = result["response"] + + contains_audio_done_event = False + for response in responses: + if response["type"] == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE.value: + contains_audio_delta = True + break + assert contains_audio_delta, "Expected audio delta event" diff --git a/tests/litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 6438fbe778..455d9fb9d1 100644 --- a/tests/litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -316,23 +316,21 @@ def test_vertex_ai_candidate_token_count_inclusive( assert usage.completion_tokens == expected_usage.completion_tokens assert usage.total_tokens == expected_usage.total_tokens + def test_streaming_chunk_includes_reasoning_tokens(): - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ModelResponseIterator + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + # Simulate a streaming chunk as would be received from Gemini chunk = { - "candidates": [ - { - "content": { - "parts": [{"text": "Hello"}] - } - } - ], + "candidates": [{"content": {"parts": [{"text": "Hello"}]}}], "usageMetadata": { "promptTokenCount": 5, "candidatesTokenCount": 7, "totalTokenCount": 12, "thoughtsTokenCount": 3, - } + }, } iterator = ModelResponseIterator(streaming_response=[], sync_stream=True) streaming_chunk = iterator.chunk_parser(chunk) @@ -340,7 +338,10 @@ def test_streaming_chunk_includes_reasoning_tokens(): assert streaming_chunk["usage"]["prompt_tokens"] == 5 assert streaming_chunk["usage"]["completion_tokens"] == 7 assert streaming_chunk["usage"]["total_tokens"] == 12 - assert streaming_chunk["usage"]["completion_tokens_details"]["reasoning_tokens"] == 3 + assert ( + streaming_chunk["usage"]["completion_tokens_details"]["reasoning_tokens"] == 3 + ) + def test_check_finish_reason(): config = VertexGeminiConfig() @@ -350,3 +351,27 @@ def test_check_finish_reason(): config._check_finish_reason(chat_completion_message=None, finish_reason=k) == v ) + + +def test_vertex_ai_usage_metadata_response_token_count(): + """For Gemini Live API""" + from litellm.types.utils import PromptTokensDetailsWrapper + + v = VertexGeminiConfig() + usage_metadata = { + "promptTokenCount": 57, + "responseTokenCount": 74, + "totalTokenCount": 131, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 57}], + "responseTokensDetails": [{"modality": "TEXT", "tokenCount": 74}], + } + usage_metadata = UsageMetadata(**usage_metadata) + result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + print("result", result) + assert result.prompt_tokens == 57 + assert result.completion_tokens == 74 + assert result.total_tokens == 131 + assert result.prompt_tokens_details.text_tokens == 57 + assert result.prompt_tokens_details.audio_tokens is None + assert result.prompt_tokens_details.cached_tokens is None + assert result.completion_tokens_details.text_tokens == 74 diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 4ff469f882..f58667e944 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -3161,3 +3161,5 @@ async def test_bedrock_max_completion_tokens(model: str): "system": [], "inferenceConfig": {"maxTokens": 10}, } + +