diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 347eef70a9e..329f2b63c20 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 db98b7e56a4..d5531a532b9 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 bb66b419b9c..4cf89accfce 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 abddc766afa..59360281ac3 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 203c5634364..cd67be3545a 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 94024dc83f9..00000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -