From fdfef04d9323ee727f8e058f79ac9f22a6743877 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Thu, 15 May 2025 22:39:09 -0700 Subject: [PATCH] Gemini Multimodal Live API support (#10841) * fix: initial commit * refactor(gemini/realtime/transformation.py): initial instrumentation * fix: fix default api base if not set * feat(gemini/): passes initial user request to backend * feat(realtime_streaming.py): support transforming message input before sending it enables gemini realtime streaming to be sent in correct format * feat: initial working commit of setup message being sent + working * fix(gemini/): initial commit supporting realtime response transformation * feat(gemini/realtime): transform session.created event correctly * fix(realtime_streaming.py): more gemini/realtime response mapping - handle new message * test(gemini/realtime/test_): add more unit tests * feat(gemini/realtime): handles consecutive deltas * feat(gemini/realtime): support openai 'response.text.done' event * feat(gemini/realtime): add openai 'response.text.done' and 'response.content_part.done' event support * fix(gemini/realtime): add openai 'response.done' event support unified realtime api support * fix: fix linting errors * fix: fix linting error * fix: fix linting error * fix: handle infinite loop * fix: fix linting error * fix: fix recursive detector * fix: fix file * fix: fix linting error * fix: fix linting error --- .../litellm_core_utils/realtime_streaming.py | 179 +++-- litellm/llms/azure/realtime/handler.py | 7 +- .../llms/base_llm/realtime/transformation.py | 75 +++ litellm/llms/custom_httpx/llm_http_handler.py | 58 ++ litellm/llms/gemini/common_utils.py | 49 +- .../llms/gemini/realtime/transformation.py | 630 ++++++++++++++++++ litellm/llms/openai/realtime/handler.py | 5 +- .../vertex_and_google_ai_studio_gemini.py | 9 +- litellm/proxy/_new_secret_config.yaml | 13 +- litellm/realtime_api/main.py | 33 +- .../transformation.py | 13 +- litellm/types/llms/gemini.py | 41 ++ litellm/types/llms/openai.py | 244 ++++++- litellm/types/realtime.py | 29 + litellm/utils.py | 12 + .../code_coverage_tests/recursive_detector.py | 3 +- .../test_gemini_realtime_transformation.py | 93 +++ 17 files changed, 1407 insertions(+), 86 deletions(-) create mode 100644 litellm/llms/base_llm/realtime/transformation.py create mode 100644 litellm/llms/gemini/realtime/transformation.py create mode 100644 litellm/types/realtime.py create mode 100644 tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 5dcabe2dd3..347eef70a9 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1,43 +1,28 @@ -""" -async with websockets.connect( # type: ignore - url, - extra_headers={ - "api-key": api_key, # type: ignore - }, - ) as backend_ws: - forward_task = asyncio.create_task( - forward_messages(websocket, backend_ws) - ) - - try: - while True: - message = await websocket.receive_text() - await backend_ws.send(message) - except websockets.exceptions.ConnectionClosed: # type: ignore - forward_task.cancel() - finally: - if not forward_task.done(): - forward_task.cancel() - try: - await forward_task - except asyncio.CancelledError: - pass -""" - import asyncio import concurrent.futures import json -from typing import Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union import litellm from litellm._logging import verbose_logger +from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.types.llms.openai import ( + OpenAIRealtimeEvents, + OpenAIRealtimeOutputItemDone, + OpenAIRealtimeResponseTextDelta, OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents, ) from .litellm_logging import Logging as LiteLLMLogging +if TYPE_CHECKING: + from websockets.asyncio.client import ClientConnection + + CLIENT_CONNECTION_CLASS = ClientConnection +else: + CLIENT_CONNECTION_CLASS = Any + # Create a thread pool with a maximum of 10 threads executor = concurrent.futures.ThreadPoolExecutor(max_workers=10) @@ -52,18 +37,15 @@ class RealTimeStreaming: def __init__( self, websocket: Any, - backend_ws: Any, - logging_obj: Optional[LiteLLMLogging] = None, + backend_ws: CLIENT_CONNECTION_CLASS, + logging_obj: LiteLLMLogging, + provider_config: Optional[BaseRealtimeConfig] = None, + model: str = "", ): self.websocket = websocket self.backend_ws = backend_ws self.logging_obj = logging_obj - self.messages: List[ - Union[ - OpenAIRealtimeStreamResponseBaseObject, - OpenAIRealtimeStreamSessionEvents, - ] - ] = [] + self.messages: List[OpenAIRealtimeEvents] = [] self.input_message: Dict = {} _logged_real_time_event_types = litellm.logged_real_time_event_types @@ -71,34 +53,43 @@ class RealTimeStreaming: if _logged_real_time_event_types is None: _logged_real_time_event_types = DefaultLoggedRealTimeEventTypes 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_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 def _should_store_message( self, - message_obj: Union[ - dict, - OpenAIRealtimeStreamSessionEvents, - OpenAIRealtimeStreamResponseBaseObject, - ], + message_obj: Union[dict, OpenAIRealtimeEvents], ) -> bool: - _msg_type = message_obj["type"] + _msg_type = message_obj["type"] if "type" in message_obj else None if self.logged_real_time_event_types == "*": return True - if _msg_type in self.logged_real_time_event_types: + if _msg_type and _msg_type in self.logged_real_time_event_types: return True return False - def store_message(self, message: Union[str, bytes]): + def store_message(self, message: Union[str, bytes, OpenAIRealtimeEvents]): """Store message in list""" if isinstance(message, bytes): message = message.decode("utf-8") - message_obj = json.loads(message) + if isinstance(message, dict): + message_obj = message + else: + message_obj = json.loads(message) try: if ( - message_obj.get("type") == "session.created" + not isinstance(message, dict) + or message_obj.get("type") == "session.created" or message_obj.get("type") == "session.updated" ): message_obj = OpenAIRealtimeStreamSessionEvents(**message_obj) # type: ignore - else: + elif not isinstance(message, dict): message_obj = OpenAIRealtimeStreamResponseBaseObject(**message_obj) # type: ignore except Exception as e: verbose_logger.debug(f"Error parsing message for logging: {e}") @@ -121,20 +112,68 @@ class RealTimeStreaming: ## SYNC LOGGING executor.submit(self.logging_obj.success_handler(self.messages)) - async def backend_to_client_send_messages(self): + async def backend_to_client_send_messages( + self, session_configuration_request: Optional[str] = None + ): import websockets try: while True: - message = await self.backend_ws.recv() - await self.websocket.send_text(message) + try: + raw_response = await self.backend_ws.recv( + decode=False + ) # improves performance + except TypeError: + raw_response = await self.backend_ws.recv() # type: ignore[assignment] - ## LOGGING - self.store_message(message) - except websockets.exceptions.ConnectionClosed: # type: ignore - pass - except Exception: - pass + if self.provider_config: + returned_object = self.provider_config.transform_realtime_response( + raw_response, + self.model, + self.logging_obj, + realtime_response_transform_input={ + "session_configuration_request": 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, + }, + ) + + transformed_response = returned_object["response"] + self.current_output_item_id = returned_object[ + "current_output_item_id" + ] + self.current_response_id = returned_object["current_response_id"] + self.current_delta_chunks = returned_object["current_delta_chunks"] + self.current_conversation_id = returned_object[ + "current_conversation_id" + ] + self.current_item_chunks = returned_object["current_item_chunks"] + if isinstance(transformed_response, list): + for event in transformed_response: + event_str = json.dumps(event) + ## LOGGING + self.store_message(event_str) + await self.websocket.send_text(event_str) + else: + event_str = json.dumps(transformed_response) + ## LOGGING + self.store_message(event_str) + await self.websocket.send_text(event_str) + + else: + ## LOGGING + self.store_message(raw_response) + await self.websocket.send_text(raw_response) + + except websockets.exceptions.ConnectionClosed as e: # type: ignore + verbose_logger.exception( + f"Connection closed in backend to client send messages - {e}" + ) + except Exception as e: + verbose_logger.exception(f"Error in backend to client send messages: {e}") finally: await self.log_messages() @@ -142,18 +181,42 @@ class RealTimeStreaming: try: while True: message = await self.websocket.receive_text() + ## LOGGING self.store_input(message=message) ## FORWARD TO BACKEND + if self.provider_config: + message = self.provider_config.transform_realtime_request(message) + await self.backend_ws.send(message) - except self.websockets.exceptions.ConnectionClosed: # type: ignore + 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): - forward_task = asyncio.create_task(self.backend_to_client_send_messages()) + 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) + ) try: await self.client_ack_messages() - except self.websockets.exceptions.ConnectionClosed: # type: ignore + except self.websocket.exceptions.ConnectionClosed: # type: ignore + verbose_logger.debug("Connection closed") forward_task.cancel() finally: if not forward_task.done(): diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index 5a4865e7d7..c5447b4ccd 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -4,7 +4,7 @@ This file contains the calling Azure OpenAI's `/openai/realtime` endpoint. This requires websockets, and is currently only supported on LiteLLM Proxy. """ -from typing import Any, Optional +from typing import Any, Optional, cast from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from ....litellm_core_utils.realtime_streaming import RealTimeStreaming @@ -40,15 +40,16 @@ class AzureOpenAIRealtime(AzureChatCompletion): self, model: str, websocket: Any, + logging_obj: LiteLLMLogging, api_base: Optional[str] = None, api_key: Optional[str] = None, api_version: Optional[str] = None, azure_ad_token: Optional[str] = None, client: Optional[Any] = None, - logging_obj: Optional[LiteLLMLogging] = None, timeout: Optional[float] = None, ): import websockets + from websockets.asyncio.client import ClientConnection if api_base is None: raise ValueError("api_base is required for Azure OpenAI calls") @@ -65,7 +66,7 @@ class AzureOpenAIRealtime(AzureChatCompletion): }, ) as backend_ws: realtime_streaming = RealTimeStreaming( - websocket, backend_ws, logging_obj + websocket, cast(ClientConnection, backend_ws), logging_obj ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/llms/base_llm/realtime/transformation.py b/litellm/llms/base_llm/realtime/transformation.py new file mode 100644 index 0000000000..759cda4744 --- /dev/null +++ b/litellm/llms/base_llm/realtime/transformation.py @@ -0,0 +1,75 @@ +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING, Any, Optional, Union + +import httpx + +from litellm.types.realtime import ( + RealtimeResponseTransformInput, + RealtimeResponseTypedDict, +) + +from ..chat.transformation import BaseLLMException + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class BaseRealtimeConfig(ABC): + @abstractmethod + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + ) -> dict: + pass + + @abstractmethod + def get_complete_url( + self, api_base: Optional[str], model: str, api_key: Optional[str] = None + ) -> str: + """ + OPTIONAL + + Get the complete url for the request + + Some providers need `model` in `api_base` + """ + return api_base or "" + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BaseLLMException: + raise BaseLLMException( + status_code=status_code, + message=error_message, + headers=headers, + ) + + @abstractmethod + def transform_realtime_request(self, message: str) -> str: + pass + + def requires_session_configuration( + self, + ) -> bool: # initial configuration message sent to setup the realtime session + return False + + def session_configuration_request( + self, model: str + ) -> Optional[str]: # message sent to setup the realtime session + return None + + @abstractmethod + def transform_realtime_response( + self, + message: Union[str, bytes], + model: str, + logging_obj: LiteLLMLoggingObj, + realtime_response_transform_input: RealtimeResponseTransformInput, + ) -> RealtimeResponseTypedDict: # message sent to setup the realtime session + pass diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index fb22a2f4ea..bb66b419b9 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -19,6 +19,7 @@ import litellm.litellm_core_utils import litellm.types import litellm.types.utils from litellm._logging import verbose_logger +from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming from litellm.llms.base_llm.anthropic_messages.transformation import ( BaseAnthropicMessagesConfig, ) @@ -29,6 +30,7 @@ from litellm.llms.base_llm.base_model_iterator import MockResponseIterator from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.llms.base_llm.files.transformation import BaseFilesConfig +from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.llms.custom_httpx.http_handler import ( @@ -2040,3 +2042,59 @@ class BaseLLMHTTPHandler: status_code=status_code, headers=error_headers, ) + + async def async_realtime( + self, + model: str, + websocket: Any, + logging_obj: LiteLLMLoggingObj, + provider_config: BaseRealtimeConfig, + headers: dict, + api_base: Optional[str] = None, + api_key: Optional[str] = None, + client: Optional[Any] = None, + timeout: Optional[float] = None, + ): + import websockets + from websockets.asyncio.client import ClientConnection + + url = provider_config.get_complete_url(api_base, model, api_key) + headers = provider_config.validate_environment( + headers=headers, + model=model, + api_key=api_key, + ) + + try: + async with websockets.connect( # type: ignore + url, additional_headers=headers + ) as backend_ws: + realtime_streaming = RealTimeStreaming( + websocket, + cast(ClientConnection, backend_ws), + logging_obj, + provider_config, + model, + ) + await realtime_streaming.bidirectional_forward() + + except websockets.exceptions.InvalidStatusCode as e: # type: ignore + verbose_logger.exception(f"Error connecting to backend: {e}") + await websocket.close(code=e.status_code, reason=str(e)) + except Exception as e: + verbose_logger.exception(f"Error connecting to backend: {e}") + try: + await websocket.close( + code=1011, reason=f"Internal server error: {str(e)}" + ) + except RuntimeError as close_error: + if "already completed" in str(close_error) or "websocket.close" in str( + close_error + ): + # The WebSocket is already closed or the response is completed, so we can ignore this error + pass + else: + # If it's a different RuntimeError, we might want to log it or handle it differently + raise Exception( + f"Unexpected error while closing WebSocket: {close_error}" + ) diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index fef41f7d58..3331f584b5 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -1,8 +1,11 @@ -from typing import List, Optional, Union +import base64 +import datetime +from typing import Dict, List, Optional, Union import httpx import litellm +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.llms.base_llm.base_utils import BaseLLMModelInfo from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.secret_managers.main import get_secret_str @@ -82,3 +85,47 @@ class GeminiModelInfo(BaseLLMModelInfo): return GeminiError( status_code=status_code, message=error_message, headers=headers ) + + +def encode_unserializable_types( + data: Dict[str, object], depth: int = 0 +) -> Dict[str, object]: + """Converts unserializable types in dict to json.dumps() compatible types. + + This function is called in models.py after calling convert_to_dict(). The + convert_to_dict() can convert pydantic object to dict. However, the input to + convert_to_dict() is dict mixed of pydantic object and nested dict(the output + of converters). So they may be bytes in the dict and they are out of + `ser_json_bytes` control in model_dump(mode='json') called in + `convert_to_dict`, as well as datetime deserialization in Pydantic json mode. + + Returns: + A dictionary with json.dumps() incompatible type (e.g. bytes datetime) + to compatible type (e.g. base64 encoded string, isoformat date string). + """ + if depth > DEFAULT_MAX_RECURSE_DEPTH: + return data + processed_data: dict[str, object] = {} + if not isinstance(data, dict): + return data + for key, value in data.items(): + if isinstance(value, bytes): + processed_data[key] = base64.urlsafe_b64encode(value).decode("ascii") + elif isinstance(value, datetime.datetime): + processed_data[key] = value.isoformat() + elif isinstance(value, dict): + processed_data[key] = encode_unserializable_types(value, depth + 1) + elif isinstance(value, list): + if all(isinstance(v, bytes) for v in value): + processed_data[key] = [ + base64.urlsafe_b64encode(v).decode("ascii") for v in value + ] + if all(isinstance(v, datetime.datetime) for v in value): + processed_data[key] = [v.isoformat() for v in value] + else: + processed_data[key] = [ + encode_unserializable_types(v, depth + 1) for v in value + ] + else: + processed_data[key] = value + return processed_data diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py new file mode 100644 index 0000000000..4d29a68c2f --- /dev/null +++ b/litellm/llms/gemini/realtime/transformation.py @@ -0,0 +1,630 @@ +""" +This file contains the transformation logic for the Gemini realtime API. +""" + +import json +import os +import uuid +from typing import Any, Dict, List, Optional, Union, cast + +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 ( + VertexGeminiConfig, +) +from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, +) +from litellm.types.llms.gemini import ( + BidiGenerateContentServerContent, + BidiGenerateContentServerMessage, +) +from litellm.types.llms.openai import ( + OpenAIRealtimeContentPartDone, + OpenAIRealtimeConversationItemCreated, + OpenAIRealtimeDoneEvent, + OpenAIRealtimeEvents, + OpenAIRealtimeOutputItemDone, + OpenAIRealtimeResponseContentPartAdded, + OpenAIRealtimeResponseDoneObject, + OpenAIRealtimeResponseTextDelta, + OpenAIRealtimeResponseTextDone, + OpenAIRealtimeStreamResponseBaseObject, + OpenAIRealtimeStreamResponseOutputItemAdded, + OpenAIRealtimeStreamSession, + OpenAIRealtimeStreamSessionEvents, +) +from litellm.types.realtime import ( + RealtimeResponseTransformInput, + RealtimeResponseTypedDict, +) + +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", +} + + +class GeminiRealtimeConfig(BaseRealtimeConfig): + def validate_environment( + self, headers: dict, model: str, api_key: Optional[str] = None + ) -> dict: + return headers + + def get_complete_url( + self, api_base: Optional[str], model: str, api_key: Optional[str] = None + ) -> str: + """ + Example output: + "BACKEND_WS_URL = "wss://generativelanguage.googleapis.com/ws/google.ai.generativelanguage.v1beta.GenerativeService.BidiGenerateContent""; + """ + if api_base is None: + api_base = "wss://generativelanguage.googleapis.com" + if api_key is None: + api_key = os.environ.get("GEMINI_API_KEY") + if api_key is None: + raise ValueError("api_key is required for Gemini API calls") + api_base = api_base.replace("https://", "wss://") + 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 + + if len(realtime_input_dict) != 1: + raise ValueError( + f"Only one argument can be set, got {len(realtime_input_dict)}:" + f" {list(realtime_input_dict.keys())}" + ) + + realtime_input_dict = encode_unserializable_types(realtime_input_dict) + + return json.dumps({"realtime_input": realtime_input_dict}) + + def transform_session_created_event( + self, + model: str, + 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" + ) + + 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"]) + _system_instruction = session_configuration_request_dict.get( + "systemInstruction" + ) + session = OpenAIRealtimeStreamSession( + id=logging_session_id, + modalities=_modalities, + ) + 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 + + return OpenAIRealtimeStreamSessionEvents( + type="session.created", + session=session, + event_id=str(uuid.uuid4()), + ) + + def _is_new_content_delta( + self, + previous_messages: Optional[List[OpenAIRealtimeEvents]] = None, + ) -> bool: + if previous_messages is None or len(previous_messages) == 0: + return True + if "type" in previous_messages[-1] and previous_messages[-1]["type"].endswith( + "delta" + ): + return False + return True + + def return_new_content_delta_events( + self, + response_id: str, + output_item_id: str, + conversation_id: str, + session_configuration_request: Optional[str] = None, + ) -> List[OpenAIRealtimeEvents]: + if session_configuration_request is None: + raise ValueError( + "session_configuration_request is required for Gemini API calls" + ) + + session_configuration_request_dict = json.loads(session_configuration_request) + _modalities = 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") + + response_items: List[OpenAIRealtimeEvents] = [] + + ## - return response.created + response_created = OpenAIRealtimeStreamResponseBaseObject( + type="response.created", + event_id="event_{}".format(uuid.uuid4()), + response={ + "object": "realtime.response", + "id": response_id, + "status": "in_progress", + "output": [], + "conversation_id": conversation_id, + "modalities": _modalities, + "temperature": _temperature, + "max_output_tokens": _max_output_tokens, + }, + ) + response_items.append(response_created) + + ## - return response.output_item.added ← adds ‘item_id’ same for all subsequent events + response_output_item_added = OpenAIRealtimeStreamResponseOutputItemAdded( + type="response.output_item.added", + response_id=response_id, + output_index=0, + item={ + "id": output_item_id, + "object": "realtime.item", + "type": "message", + "status": "in_progress", + "role": "assistant", + "content": [], + }, + ) + response_items.append(response_output_item_added) + ## - return conversation.item.created + conversation_item_created = OpenAIRealtimeConversationItemCreated( + type="conversation.item.created", + event_id="event_{}".format(uuid.uuid4()), + item={ + "id": output_item_id, + "object": "realtime.item", + "type": "message", + "status": "in_progress", + "role": "assistant", + "content": [], + }, + ) + response_items.append(conversation_item_created) + ## - return response.content_part.added + response_content_part_added = OpenAIRealtimeResponseContentPartAdded( + type="response.content_part.added", + content_index=0, + output_index=0, + event_id="event_{}".format(uuid.uuid4()), + item_id=output_item_id, + part={ + "type": "text", + "text": "", + }, + response_id=response_id, + ) + response_items.append(response_content_part_added) + return response_items + + def transform_content_delta_events( + self, + message: BidiGenerateContentServerContent, + output_item_id: str, + response_id: str, + ) -> OpenAIRealtimeResponseTextDelta: + delta = "" + try: + if "modelTurn" in message and "parts" in message["modelTurn"]: + for part in message["modelTurn"]["parts"]: + if "text" in part: + delta += part["text"] + except Exception as e: + raise ValueError( + f"Error transforming content delta events: {e}, got message: {message}" + ) + + return OpenAIRealtimeResponseTextDelta( + type="response.text.delta", + content_index=0, + event_id="event_{}".format(uuid.uuid4()), + item_id=output_item_id, + output_index=0, + response_id=response_id, + delta=delta, + ) + + def transform_content_done_event( + self, + delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]], + current_output_item_id: Optional[str], + current_response_id: Optional[str], + ) -> OpenAIRealtimeResponseTextDone: + if delta_chunks: + delta = "".join([delta_chunk["delta"] for delta_chunk in delta_chunks]) + else: + delta = "" + if current_output_item_id is None or current_response_id is None: + 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, + ) + + def return_additional_content_done_events( + self, + current_output_item_id: Optional[str], + current_response_id: Optional[str], + delta_done_event: OpenAIRealtimeResponseTextDone, + ) -> List[OpenAIRealtimeEvents]: + """ + - return response.content_part.done + - return response.output_item.done + """ + if current_output_item_id is None or current_response_id is None: + raise ValueError( + "current_output_item_id and current_response_id cannot be None for a 'done' event." + ) + returned_items: List[OpenAIRealtimeEvents] = [] + # response.content_part.done + response_content_part_done = OpenAIRealtimeContentPartDone( + type="response.content_part.done", + content_index=0, + event_id="event_{}".format(uuid.uuid4()), + item_id=current_output_item_id, + output_index=0, + part={ + "type": "text", + "text": delta_done_event["text"], + }, + response_id=current_response_id, + ) + returned_items.append(response_content_part_done) + # response.output_item.done + response_output_item_done = OpenAIRealtimeOutputItemDone( + type="response.output_item.done", + event_id="event_{}".format(uuid.uuid4()), + output_index=0, + response_id=current_response_id, + item={ + "id": current_output_item_id, + "object": "realtime.item", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "text", + "text": delta_done_event["text"], + } + ], + }, + ) + returned_items.append(response_output_item_done) + return returned_items + + @staticmethod + def get_nested_value(obj: dict, path: str) -> Any: + keys = path.split(".") + current = obj + for key in keys: + if isinstance(current, dict) and key in current: + current = current[key] + else: + return None + return current + + def update_current_delta_chunks( + self, + transformed_message: Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]], + current_delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]], + ) -> Optional[List[OpenAIRealtimeResponseTextDelta]]: + try: + if isinstance(transformed_message, list): + current_delta_chunks = [] + any_delta_chunk = False + for event in transformed_message: + if event["type"] == "response.text.delta": + current_delta_chunks.append( + cast(OpenAIRealtimeResponseTextDelta, event) + ) + any_delta_chunk = True + if not any_delta_chunk: + current_delta_chunks = ( + None # reset current_delta_chunks if no delta chunks + ) + else: + if transformed_message["type"] == "response.text.delta": + if current_delta_chunks is None: + current_delta_chunks = [] + current_delta_chunks.append( + cast(OpenAIRealtimeResponseTextDelta, transformed_message) + ) + else: + current_delta_chunks = None + return current_delta_chunks + except Exception as e: + raise ValueError( + f"Error updating current delta chunks: {e}, got transformed_message: {transformed_message}" + ) + + def update_current_item_chunks( + self, + transformed_message: Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]], + current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]], + ) -> Optional[List[OpenAIRealtimeOutputItemDone]]: + try: + if isinstance(transformed_message, list): + current_item_chunks = [] + any_item_chunk = False + for event in transformed_message: + if event["type"] == "response.output_item.done": + current_item_chunks.append( + cast(OpenAIRealtimeOutputItemDone, event) + ) + any_item_chunk = True + if not any_item_chunk: + current_item_chunks = ( + None # reset current_item_chunks if no item chunks + ) + else: + if transformed_message["type"] == "response.output_item.done": + if current_item_chunks is None: + current_item_chunks = [] + current_item_chunks.append( + cast(OpenAIRealtimeOutputItemDone, transformed_message) + ) + else: + current_item_chunks = None + return current_item_chunks + except Exception as e: + raise ValueError( + f"Error updating current item chunks: {e}, got transformed_message: {transformed_message}" + ) + + def transform_response_done_event( + self, + 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 + ): + 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" + ) + + session_configuration_request_dict = json.loads(session_configuration_request) + temperature = 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, + ) + responses_api_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + _chat_completion_usage, + ) + return OpenAIRealtimeDoneEvent( + type="response.done", + event_id="event_{}".format(uuid.uuid4()), + response=OpenAIRealtimeResponseDoneObject( + object="realtime.response", + id=current_response_id, + status="completed", + output=[output_item["item"] for output_item in output_items] + if output_items + else [], + conversation_id=current_conversation_id, + modalities=_modalities, + temperature=temperature, + max_output_tokens=max_output_tokens, + usage=responses_api_usage.model_dump(), + ), + ) + + def transform_realtime_response( + self, + message: Union[str, bytes], + model: str, + logging_obj: LiteLLMLoggingObj, + realtime_response_transform_input: RealtimeResponseTransformInput, + ) -> RealtimeResponseTypedDict: + 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}") + + logging_session_id = logging_obj.litellm_trace_id + 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" + ] + current_item_chunks = realtime_response_transform_input["current_item_chunks"] + returned_message: Optional[ + Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]] + ] = None + 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 + + 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 isinstance(message, bytes): + message_str = message.decode("utf-8", errors="replace") + else: + message_str = str(message) + raise ValueError(f"Unknown message type: {message_str}") + + current_delta_chunks = self.update_current_delta_chunks( + transformed_message=returned_message, + current_delta_chunks=current_delta_chunks, + ) + current_item_chunks = self.update_current_item_chunks( + transformed_message=returned_message, + current_item_chunks=current_item_chunks, + ) + return { + "response": returned_message, + "current_output_item_id": current_output_item_id, + "current_response_id": current_response_id, + "current_delta_chunks": current_delta_chunks, + "current_conversation_id": current_conversation_id, + "current_item_chunks": current_item_chunks, + } + + def requires_session_configuration(self) -> bool: + return True + + def session_configuration_request(self, model: str) -> Optional[str]: + """ + + ``` + { + "model": string, + "generationConfig": { + "candidateCount": integer, + "maxOutputTokens": integer, + "temperature": number, + "topP": number, + "topK": integer, + "presencePenalty": number, + "frequencyPenalty": number, + "responseModalities": [string], + "speechConfig": object, + "mediaResolution": object + }, + "systemInstruction": string, + "tools": [object] + } + ``` + """ + return json.dumps( + { + "setup": { + "model": f"models/{model}", + "generationConfig": {"responseModalities": ["TEXT"]}, + } + } + ) diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index 83398ad11a..099eeab7e5 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -4,7 +4,7 @@ This file contains the calling Azure OpenAI's `/openai/realtime` endpoint. This requires websockets, and is currently only supported on LiteLLM Proxy. """ -from typing import Any, Optional +from typing import Any, Optional, cast from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from ....litellm_core_utils.realtime_streaming import RealTimeStreaming @@ -32,6 +32,7 @@ class OpenAIRealtime(OpenAIChatCompletion): timeout: Optional[float] = None, ): import websockets + from websockets.asyncio.client import ClientConnection if api_base is None: raise ValueError("api_base is required for Azure OpenAI calls") @@ -49,7 +50,7 @@ class OpenAIRealtime(OpenAIChatCompletion): }, ) as backend_ws: realtime_streaming = RealTimeStreaming( - websocket, backend_ws, logging_obj + websocket, cast(ClientConnection, backend_ws), logging_obj ) await realtime_streaming.bidirectional_forward() 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 c230d062cc..203c563436 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 @@ -36,6 +36,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, ) from litellm.types.llms.anthropic import AnthropicThinkingParam +from litellm.types.llms.gemini import BidiGenerateContentServerMessage from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionResponseMessage, @@ -789,8 +790,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def _calculate_usage( self, - completion_response: GenerateContentResponseBody, + completion_response: Union[ + GenerateContentResponseBody, BidiGenerateContentServerMessage + ], ) -> Usage: + if "usageMetadata" not in completion_response: + raise ValueError( + f"usageMetadata not found in completion_response. Got={completion_response}" + ) cached_tokens: Optional[int] = None audio_tokens: Optional[int] = None text_tokens: Optional[int] = None diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index b7d1f4188f..6ddb8c5d06 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,9 +1,7 @@ model_list: - model_name: "gemini-2.0-flash" litellm_params: - model: vertex_ai/gemini-2.0-flash - vertex_project: my-project-id - vertex_location: us-central1 + model: gemini/gemini-2.0-flash-live-001 - model_name: "gpt-4o-mini-openai" litellm_params: model: gpt-4o-mini @@ -29,11 +27,9 @@ model_list: model: databricks/databricks-claude-3-7-sonnet api_key: os.environ/DATABRICKS_API_KEY api_base: os.environ/DATABRICKS_API_BASE - - model_name: "gpt-4.1" + - model_name: gpt-4.1 litellm_params: - model: azure/gpt-4.1 - api_key: os.environ/AZURE_API_KEY_REALTIME - api_base: https://krris-m2f9a9i7-eastus2.openai.azure.com/ + model: openai/gpt-4o-realtime-preview - model_name: "xai/*" litellm_params: model: xai/* @@ -72,6 +68,3 @@ model_list: model_info: id: my-unique-azure-deployment mode: batch - -litellm_settings: - success_callback: ["generic_api"] diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 42203b46f8..fcf6c21845 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -1,11 +1,15 @@ """Abstraction function for OpenAI's realtime API""" -from typing import Any, Optional +from typing import Any, Optional, cast import litellm from litellm import get_llm_provider +from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager from ..litellm_core_utils.get_litellm_params import get_litellm_params from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging @@ -15,6 +19,7 @@ from ..utils import client as wrapper_client azure_realtime = AzureOpenAIRealtime() openai_realtime = OpenAIRealtime() +base_llm_http_handler = BaseLLMHTTPHandler() @wrapper_client @@ -34,6 +39,12 @@ async def _arealtime( For PROXY use only. """ + headers = cast(Optional[dict], kwargs.get("headers")) + extra_headers = cast(Optional[dict], kwargs.get("extra_headers")) + if headers is None: + headers = {} + if extra_headers is not None: + headers.update(extra_headers) litellm_logging_obj: LiteLLMLogging = kwargs.get("litellm_logging_obj") # type: ignore user = kwargs.get("user", None) litellm_params = GenericLiteLLMParams(**kwargs) @@ -54,7 +65,25 @@ async def _arealtime( custom_llm_provider=_custom_llm_provider, ) - if _custom_llm_provider == "azure": + provider_config: Optional[BaseRealtimeConfig] = None + if _custom_llm_provider in LlmProviders._member_map_.values(): + provider_config = ProviderConfigManager.get_provider_realtime_config( + model=model, + provider=LlmProviders(_custom_llm_provider), + ) + if provider_config is not None: + await base_llm_http_handler.async_realtime( + model=model, + websocket=websocket, + logging_obj=litellm_logging_obj, + provider_config=provider_config, + api_base=api_base, + api_key=api_key, + client=client, + timeout=timeout, + headers=headers, + ) + elif _custom_llm_provider == "azure": api_base = ( dynamic_api_base or litellm_params.api_base diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 31058d439c..5812daad46 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -211,9 +211,9 @@ class LiteLLMCompletionResponsesConfig: _messages = litellm_completion_request.get("messages") or [] session_messages = chat_completion_session.get("messages") or [] litellm_completion_request["messages"] = session_messages + _messages - litellm_completion_request["litellm_trace_id"] = ( - chat_completion_session.get("litellm_session_id") - ) + litellm_completion_request[ + "litellm_trace_id" + ] = chat_completion_session.get("litellm_session_id") return litellm_completion_request @staticmethod @@ -671,9 +671,12 @@ class LiteLLMCompletionResponsesConfig: @staticmethod def _transform_chat_completion_usage_to_responses_usage( - chat_completion_response: ModelResponse, + chat_completion_response: Union[ModelResponse, Usage], ) -> ResponseAPIUsage: - usage: Optional[Usage] = getattr(chat_completion_response, "usage", None) + if isinstance(chat_completion_response, ModelResponse): + usage: Optional[Usage] = getattr(chat_completion_response, "usage", None) + else: + usage = chat_completion_response if usage is None: return ResponseAPIUsage( input_tokens=0, diff --git a/litellm/types/llms/gemini.py b/litellm/types/llms/gemini.py index ae15bd07b5..c381440881 100644 --- a/litellm/types/llms/gemini.py +++ b/litellm/types/llms/gemini.py @@ -3,6 +3,8 @@ from typing import Any, Dict, Iterable, List, Literal, Optional, Union from typing_extensions import Required, TypedDict +from .vertex_ai import HttpxContentType, UsageMetadata + class GeminiFilesState(Enum): STATE_UNSPECIFIED = "STATE_UNSPECIFIED" @@ -31,3 +33,42 @@ class GeminiCreateFilesResponseObject(TypedDict): source: GeminiFilesSource error: dict metadata: dict + + +class BidiGenerateContentTranscription(TypedDict): + text: str + """Output only. The transcription of the audio.""" + + +class BidiGenerateContentServerContent(TypedDict, total=False): + generationComplete: bool + """Output only. If true, indicates that the model is done generating.""" + + turnComplete: bool + """Output only. If true, indicates that the model has completed its turn. Generation will only start in response to additional client messages.""" + + interrupted: bool + """Output only. If true, indicates that a client message has interrupted current model generation. If the client is playing out the content in real time, this is a good signal to stop and empty the current playback queue.""" + + groundingMetadata: dict + """Output only. Grounding metadata for the generated content.""" + + inputTranscription: BidiGenerateContentTranscription + """Output only. Input audio transcription. The transcription is sent independently of the other server messages and there is no guaranteed ordering.""" + + outputTranscription: BidiGenerateContentTranscription + """Output only. Output audio transcription. The transcription is sent independently of the other server messages and there is no guaranteed ordering, in particular not between serverContent and this outputTranscription.""" + + modelTurn: HttpxContentType + """Output only. The content that the model is currently generating.""" + + +class BidiGenerateContentServerMessage(TypedDict, total=False): + usageMetadata: UsageMetadata + """Output only. Usage metadata for the generated content.""" + + serverContent: BidiGenerateContentServerContent + """Output only. The content that the model is currently generating.""" + + setupComplete: dict + """Output only. The setup complete message.""" diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index f33dbdca2c..8ae059d706 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1265,22 +1265,260 @@ ResponsesAPIStreamingResponse = Annotated[ REASONING_EFFORT = Literal["low", "medium", "high"] +class OpenAIRealtimeStreamSession(TypedDict, total=False): + id: Required[str] + """ + Unique identifier for the session that looks like sess_1234567890abcdef. + """ + + input_audio_format: str + """ + The format of input audio. Options are pcm16, g711_ulaw, or g711_alaw. For pcm16, input audio must be 16-bit PCM at a 24kHz sample rate, single channel (mono), and little-endian byte order. + """ + + input_audio_noise_reduction: object + """ + Configuration for input audio noise reduction. This can be set to null to turn off. Noise reduction filters audio added to the input audio buffer before it is sent to VAD and the model. Filtering the audio can improve VAD and turn detection accuracy (reducing false positives) and model performance by improving perception of the input audio. + """ + + input_audio_transcription: object + """ + Configuration for input audio transcription, defaults to off and can be set to null to turn off once on. Input audio transcription is not native to the model, since the model consumes audio directly. Transcription runs asynchronously through the /audio/transcriptions endpoint and should be treated as guidance of input audio content rather than precisely what the model heard. The client can optionally set the language and prompt for transcription, these offer additional guidance to the transcription service. + """ + + instructions: str + """ + The default system instructions (i.e. system message) prepended to model calls. This field allows the client to guide the model on desired responses. The model can be instructed on response content and format, (e.g. "be extremely succinct", "act friendly", "here are examples of good responses") and on audio behavior (e.g. "talk quickly", "inject emotion into your voice", "laugh frequently"). The instructions are not guaranteed to be followed by the model, but they provide guidance to the model on the desired behavior. + """ + + max_response_output_tokens: Union[int, Literal["inf"]] + """ + Maximum number of output tokens for a single assistant response, inclusive of tool calls. Provide an integer between 1 and 4096 to limit output tokens, or inf for the maximum available tokens for a given model. Defaults to inf. + """ + + modalities: List[str] + """ + The set of modalities the model can respond with. To disable audio, set this to ["text"]. + """ + + model: str + """ + The Realtime model used for this session. + """ + + output_audio_format: str + """ + The format of output audio. Options are pcm16, g711_ulaw, or g711_alaw. For pcm16, output audio is sampled at a rate of 24kHz. + """ + + temperature: float + """ + Sampling temperature for the model, limited to [0.6, 1.2]. For audio models a temperature of 0.8 is highly recommended for best performance. + """ + + tool_choice: str + """ + How the model chooses tools. Options are auto, none, required, or specify a function. + """ + + tools: list + """ + Tools (functions) available to the model. + """ + + turn_detection: object + """ + + Configuration for turn detection, ether Server VAD or Semantic VAD. This can be set to null to turn off, in which case the client must manually trigger model response. Server VAD means that the model will detect the start and end of speech based on audio volume and respond at the end of user speech. Semantic VAD is more advanced and uses a turn detection model (in conjuction with VAD) to semantically estimate whether the user has finished speaking, then dynamically sets a timeout based on this probability. For example, if user audio trails off with "uhhm", the model will score a low probability of turn end and wait longer for the user to continue speaking. This can be useful for more natural conversations, but may have a higher latency. + """ + + voice: str + """ + The voice the model uses to respond. + """ + + class OpenAIRealtimeStreamSessionEvents(TypedDict): event_id: str - session: dict + session: OpenAIRealtimeStreamSession type: Union[Literal["session.created"], Literal["session.updated"]] +class OpenAIRealtimeStreamResponseOutputItemContent(TypedDict, total=False): + audio: str + """Base64-encoded audio bytes, used for 'input_audio' content types""" + id: str + """The ID of the previous conversation item for reference""" + text: str + """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"] + """The type of content""" + + +class OpenAIRealtimeStreamResponseOutputItem(TypedDict, total=False): + arguments: str + """For function call items""" + + call_id: str + """The ID of the function call""" + + id: str + """The ID of the previous conversation item for reference""" + + content: List[OpenAIRealtimeStreamResponseOutputItemContent] + + name: str + """The name of the function call""" + + object: Literal["realtime.item"] + """The object type""" + + role: Literal["assistant", "user", "system"] + """The role of the item, only used for 'message' items""" + + status: Literal["completed", "incomplete", "in_progress"] + """The status of the item""" + + output: str + """The output of the function call""" + + type: Literal["function_call", "message", "function_call_output"] + """The type of item""" + + +class OpenAIRealtimeStreamResponseOutputItemAdded(TypedDict): + type: Literal["response.output_item.added"] + response_id: str + output_index: int + item: OpenAIRealtimeStreamResponseOutputItem + + class OpenAIRealtimeStreamResponseBaseObject(TypedDict): event_id: str response: dict type: str -OpenAIRealtimeStreamList = List[ - Union[OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents] +class OpenAIRealtimeConversationObject(TypedDict, total=False): + id: str + object: Required[Literal["realtime.conversation"]] + + +class OpenAIRealtimeConversationCreated(TypedDict, total=False): + type: Required[Literal["conversation.created"]] + conversation: OpenAIRealtimeConversationObject + event_id: str + + +class OpenAIRealtimeConversationItemCreated(TypedDict, total=False): + type: Required[Literal["conversation.item.created"]] + item: OpenAIRealtimeStreamResponseOutputItem + event_id: str + previous_item_id: str + + +class OpenAIRealtimeResponseContentPart(TypedDict, total=False): + audio: str + """Base64-encoded audio bytes, if type is 'audio'""" + + text: str + """The text content, if type is 'text'""" + + transcript: str + """The transcript content, if type is 'audio'""" + + type: Literal["audio", "text"] + """The type of content""" + + +class OpenAIRealtimeResponseContentPartAdded(TypedDict): + type: Literal["response.content_part.added"] + content_index: int + event_id: str + item_id: str + output_index: int + part: OpenAIRealtimeResponseContentPart + response_id: str + + +class OpenAIRealtimeResponseTextDelta(TypedDict): + content_index: int + delta: str + event_id: str + item_id: str + output_index: int + response_id: str + type: Literal["response.text.delta"] + + +class OpenAIRealtimeResponseTextDone(TypedDict): + content_index: int + event_id: str + item_id: str + output_index: int + response_id: str + text: str + type: Literal["response.text.done"] + + +class OpenAIRealtimeContentPartDone(TypedDict): + content_index: int + event_id: str + item_id: str + output_index: int + response_id: str + part: OpenAIRealtimeResponseContentPart + type: Literal["response.content_part.done"] + + +class OpenAIRealtimeOutputItemDone(TypedDict): + event_id: str + item: OpenAIRealtimeStreamResponseOutputItem + output_index: int + response_id: str + type: Literal["response.output_item.done"] + + +class OpenAIRealtimeResponseDoneObject(TypedDict, total=False): + conversation_id: str + id: str + max_output_tokens: int + metadata: dict + modalities: list + object: Literal["realtime.response"] + output: List[OpenAIRealtimeStreamResponseOutputItem] + output_audio_format: str + status: Literal["completed", "cancelled", "failed", "incomplete"] + status_details: dict + temperature: float + usage: dict # ResponseAPIUsage + voice: str + + +class OpenAIRealtimeDoneEvent(TypedDict): + event_id: str + response: OpenAIRealtimeResponseDoneObject + type: Literal["response.done"] + + +OpenAIRealtimeEvents = Union[ + OpenAIRealtimeStreamResponseBaseObject, + OpenAIRealtimeStreamSessionEvents, + OpenAIRealtimeStreamResponseOutputItemAdded, + OpenAIRealtimeResponseContentPartAdded, + OpenAIRealtimeConversationItemCreated, + OpenAIRealtimeConversationCreated, + OpenAIRealtimeResponseTextDelta, + OpenAIRealtimeResponseTextDone, + OpenAIRealtimeContentPartDone, + OpenAIRealtimeOutputItemDone, + OpenAIRealtimeDoneEvent, ] +OpenAIRealtimeStreamList = List[OpenAIRealtimeEvents] + class ImageGenerationRequestQuality(str, Enum): LOW = "low" diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py new file mode 100644 index 0000000000..f8d613f7f0 --- /dev/null +++ b/litellm/types/realtime.py @@ -0,0 +1,29 @@ +from typing import List, Optional, TypedDict, Union + +from .llms.openai import ( + OpenAIRealtimeEvents, + OpenAIRealtimeOutputItemDone, + OpenAIRealtimeResponseTextDelta, +) + + +class RealtimeResponseTransformInput(TypedDict): + session_configuration_request: Optional[str] + current_output_item_id: Optional[ + str + ] # used to check if this is a new content.delta or a continuation of a previous content.delta + 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_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]] + current_conversation_id: Optional[str] + + +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_conversation_id: Optional[str] + current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]] diff --git a/litellm/utils.py b/litellm/utils.py index fbae31d41f..3e9a9bdfdf 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -235,6 +235,7 @@ from litellm.llms.base_llm.image_generation.transformation import ( from litellm.llms.base_llm.image_variations.transformation import ( BaseImageVariationConfig, ) +from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig @@ -6630,6 +6631,17 @@ class ProviderConfigManager: return get_azure_image_generation_config(model) return None + @staticmethod + def get_provider_realtime_config( + model: str, + provider: LlmProviders, + ) -> Optional[BaseRealtimeConfig]: + if LlmProviders.GEMINI == provider: + from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig + + return GeminiRealtimeConfig() + return None + def get_end_user_id_for_cost_tracking( litellm_params: dict, diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 81a46535e6..c2f0d7a1a9 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -20,7 +20,8 @@ IGNORE_FUNCTIONS = [ "_sanitize_value", # testing added for circular reference "set_schema_property_ordering", # testing added for infinite recursion "process_items", # testing added for infinite recursion + max depth set. - "_can_object_call_model" # # max depth set. + "_can_object_call_model", # max depth set. + "encode_unserializable_types" # max depth set. ] diff --git a/tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py b/tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py new file mode 100644 index 0000000000..4a7ef4d85d --- /dev/null +++ b/tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py @@ -0,0 +1,93 @@ +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig +from litellm.types.llms.openai import OpenAIRealtimeStreamSessionEvents + + +def test_gemini_realtime_transformation_session_created(): + config = GeminiRealtimeConfig() + assert config is not None + + session_configuration_request = { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["TEXT"]}, + } + session_configuration_request_str = json.dumps(session_configuration_request) + session_created_message = {"setupComplete": {}} + + session_created_message_str = json.dumps(session_created_message) + logging_obj = MagicMock() + logging_obj.litellm_trace_id.return_value = "123" + + transformed_message = config.transform_realtime_response( + session_created_message_str, + "gemini-1.5-flash", + logging_obj, + session_configuration_request_str, + ) + assert transformed_message["response"]["type"] == "session.created" + + +def test_gemini_realtime_transformation_content_delta(): + config = GeminiRealtimeConfig() + assert config is not None + + session_configuration_request = { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["TEXT"]}, + } + session_configuration_request_str = json.dumps(session_configuration_request) + session_created_message = { + "serverContent": { + "modelTurn": { + "parts": [ + {"text": "Hello, world!"}, + {"text": "How are you?"}, + ] + } + } + } + + session_created_message_str = json.dumps(session_created_message) + logging_obj = MagicMock() + logging_obj.litellm_trace_id.return_value = "123" + + returned_object = config.transform_realtime_response( + session_created_message_str, + "gemini-1.5-flash", + logging_obj, + session_configuration_request_str, + ) + transformed_message = returned_object["response"] + assert isinstance(transformed_message, list) + print(transformed_message) + transformed_message_str = json.dumps(transformed_message) + assert "Hello, world" in transformed_message_str + assert "How are you?" in transformed_message_str + print(transformed_message) + + ## assert all instances of 'event_id' are unique + event_ids = [ + event["event_id"] for event in transformed_message if "event_id" in event + ] + assert len(event_ids) == len(set(event_ids)) + ## assert all instances of 'response_id' are the same + response_ids = [ + event["response_id"] for event in transformed_message if "response_id" in event + ] + assert len(set(response_ids)) == 1 + ## assert all instances of 'output_item_id' are the same + output_item_ids = [ + event["item_id"] for event in transformed_message if "item_id" in event + ] + assert len(set(output_item_ids)) == 1