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
This commit is contained in:
Krish Dholakia
2025-05-15 22:39:09 -07:00
committed by GitHub
parent ffcdf441d2
commit fdfef04d93
17 changed files with 1407 additions and 86 deletions
+121 -58
View File
@@ -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():
+4 -3
View File
@@ -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()
@@ -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
@@ -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}"
)
+48 -1
View File
@@ -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
@@ -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"]},
}
}
)
+3 -2
View File
@@ -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()
@@ -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
+3 -10
View File
@@ -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"]
+31 -2
View File
@@ -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
@@ -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,
+41
View File
@@ -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."""
+241 -3
View File
@@ -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"
+29
View File
@@ -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]]
+12
View File
@@ -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,
@@ -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.
]
@@ -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