Krrish Dholakia 2025-05-17 07:36:56 -07:00
parent cc626ad3ec
commit 5c5699b65d
15 changed files with 830 additions and 224 deletions

View file

@ -9,10 +9,11 @@ from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.types.llms.openai import (
OpenAIRealtimeEvents,
OpenAIRealtimeOutputItemDone,
OpenAIRealtimeResponseTextDelta,
OpenAIRealtimeResponseDelta,
OpenAIRealtimeStreamResponseBaseObject,
OpenAIRealtimeStreamSessionEvents,
)
from litellm.types.realtime import ALL_DELTA_TYPES
from .litellm_logging import Logging as LiteLLMLogging
@ -55,13 +56,13 @@ class RealTimeStreaming:
self.logged_real_time_event_types = _logged_real_time_event_types
self.provider_config = provider_config
self.model = model
self.current_delta_chunks: Optional[
List[OpenAIRealtimeResponseTextDelta]
] = None
self.current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]] = None
self.current_output_item_id: Optional[str] = None
self.current_response_id: Optional[str] = None
self.current_conversation_id: Optional[str] = None
self.current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]] = None
self.current_delta_type: Optional[ALL_DELTA_TYPES] = None
self.session_configuration_request: Optional[str] = None
def _should_store_message(
self,
@ -112,9 +113,7 @@ class RealTimeStreaming:
## SYNC LOGGING
executor.submit(self.logging_obj.success_handler(self.messages))
async def backend_to_client_send_messages(
self, session_configuration_request: Optional[str] = None
):
async def backend_to_client_send_messages(self):
import websockets
try:
@ -132,12 +131,13 @@ class RealTimeStreaming:
self.model,
self.logging_obj,
realtime_response_transform_input={
"session_configuration_request": session_configuration_request,
"session_configuration_request": self.session_configuration_request,
"current_output_item_id": self.current_output_item_id,
"current_response_id": self.current_response_id,
"current_delta_chunks": self.current_delta_chunks,
"current_conversation_id": self.current_conversation_id,
"current_item_chunks": self.current_item_chunks,
"current_delta_type": self.current_delta_type,
},
)
@ -151,6 +151,10 @@ class RealTimeStreaming:
"current_conversation_id"
]
self.current_item_chunks = returned_object["current_item_chunks"]
self.current_delta_type = returned_object["current_delta_type"]
self.session_configuration_request = returned_object[
"session_configuration_request"
]
if isinstance(transformed_response, list):
for event in transformed_response:
event_str = json.dumps(event)
@ -186,33 +190,20 @@ class RealTimeStreaming:
self.store_input(message=message)
## FORWARD TO BACKEND
if self.provider_config:
message = self.provider_config.transform_realtime_request(message)
message = self.provider_config.transform_realtime_request(
message, self.model
)
for msg in message:
await self.backend_ws.send(msg)
else:
await self.backend_ws.send(message)
await self.backend_ws.send(message)
except self.websocket.exceptions.ConnectionClosed: # type: ignore
verbose_logger.debug("Connection closed")
pass
except Exception as e:
verbose_logger.debug(f"Error in client ack messages: {e}")
async def bidirectional_forward(self):
session_configuration_request: Optional[str] = None
if (
self.provider_config
and self.provider_config.requires_session_configuration()
):
session_configuration_request = (
self.provider_config.session_configuration_request(self.model)
)
if session_configuration_request is None:
raise ValueError(
"Session configuration request is None, but requires_session_configuration is True"
)
await self.backend_ws.send(session_configuration_request)
forward_task = asyncio.create_task(
self.backend_to_client_send_messages(session_configuration_request)
)
forward_task = asyncio.create_task(self.backend_to_client_send_messages())
try:
await self.client_ack_messages()
except self.websocket.exceptions.ConnectionClosed: # type: ignore

View file

@ -1,5 +1,5 @@
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Optional, Union
from typing import TYPE_CHECKING, Any, List, Optional, Union
import httpx
@ -51,7 +51,12 @@ class BaseRealtimeConfig(ABC):
)
@abstractmethod
def transform_realtime_request(self, message: str) -> str:
def transform_realtime_request(
self,
message: str,
model: str,
session_configuration_request: Optional[str] = None,
) -> List[str]:
pass
def requires_session_configuration(

View file

@ -2067,7 +2067,7 @@ class BaseLLMHTTPHandler:
try:
async with websockets.connect( # type: ignore
url, additional_headers=headers
url, extra_headers=headers
) as backend_ws:
realtime_streaming = RealTimeStreaming(
websocket,

View file

@ -7,6 +7,7 @@ import os
import uuid
from typing import Any, Dict, List, Optional, Union, cast
from litellm import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
@ -16,36 +17,51 @@ from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
from litellm.types.llms.gemini import (
AutomaticActivityDetection,
BidiGenerateContentRealtimeInput,
BidiGenerateContentRealtimeInputConfig,
BidiGenerateContentServerContent,
BidiGenerateContentServerMessage,
BidiGenerateContentSetup,
)
from litellm.types.llms.openai import (
OpenAIRealtimeContentPartDone,
OpenAIRealtimeConversationItemCreated,
OpenAIRealtimeDoneEvent,
OpenAIRealtimeEvents,
OpenAIRealtimeEventTypes,
OpenAIRealtimeOutputItemDone,
OpenAIRealtimeResponseAudioDone,
OpenAIRealtimeResponseContentPartAdded,
OpenAIRealtimeResponseDelta,
OpenAIRealtimeResponseDoneObject,
OpenAIRealtimeResponseTextDelta,
OpenAIRealtimeResponseTextDone,
OpenAIRealtimeStreamResponseBaseObject,
OpenAIRealtimeStreamResponseOutputItemAdded,
OpenAIRealtimeStreamSession,
OpenAIRealtimeStreamSessionEvents,
OpenAIRealtimeTurnDetection,
)
from litellm.types.llms.vertex_ai import (
GeminiResponseModalities,
HttpxBlobType,
HttpxContentType,
)
from litellm.types.realtime import (
ALL_DELTA_TYPES,
RealtimeModalityResponseTransformOutput,
RealtimeResponseTransformInput,
RealtimeResponseTypedDict,
)
from litellm.utils import get_empty_usage
from ..common_utils import encode_unserializable_types
MAP_GEMINI_FIELD_TO_OPENAI_EVENT = {
"setupComplete": "session.created",
"serverContent.modelTurn": "response.text.delta",
"serverContent.generationComplete": "response.text.done",
"serverContent.turnComplete": "response.done",
MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[str, OpenAIRealtimeEventTypes] = {
"setupComplete": OpenAIRealtimeEventTypes.SESSION_CREATED,
"serverContent.generationComplete": OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE,
"serverContent.turnComplete": OpenAIRealtimeEventTypes.RESPONSE_DONE,
"serverContent.interrupted": OpenAIRealtimeEventTypes.RESPONSE_DONE,
}
@ -72,9 +88,162 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
api_base = api_base.replace("http://", "ws://")
return f"{api_base}/ws/google.ai.generativelanguage.v1beta.GenerativeService.BidiGenerateContent?key={api_key}"
def transform_realtime_request(self, message: str) -> str:
realtime_input_dict: Dict[str, Any] = {}
realtime_input_dict["text"] = message
def map_model_turn_event(
self, model_turn: HttpxContentType
) -> OpenAIRealtimeEventTypes:
"""
Map the model turn event to the OpenAI realtime events.
Returns either:
- response.text.delta - model_turn: {"parts": [{"text": "..."}]}
- response.audio.delta - model_turn: {"parts": [{"inlineData": {"mimeType": "audio/pcm", "data": "..."}}]}
Assumes parts is a single element list.
"""
if "parts" in model_turn:
parts = model_turn["parts"]
if len(parts) != 1:
verbose_logger.warning(
f"Realtime: Expected 1 part, got {len(parts)} for Gemini model turn event."
)
part = parts[0]
if "text" in part:
return OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA
elif "inlineData" in part:
return OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA
else:
raise ValueError(f"Unexpected part type: {part}")
raise ValueError(f"Unexpected model turn event, no 'parts' key: {model_turn}")
def map_generation_complete_event(
self, delta_type: Optional[ALL_DELTA_TYPES]
) -> OpenAIRealtimeEventTypes:
if delta_type == "text":
return OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE
elif delta_type == "audio":
return OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE
else:
raise ValueError(f"Unexpected delta type: {delta_type}")
def get_audio_mime_type(self, input_audio_format: str = "pcm16"):
mime_types = {
"pcm16": "audio/pcm",
"g711_ulaw": "audio/pcmu",
"g711_alaw": "audio/pcma",
}
return mime_types.get(input_audio_format, "application/octet-stream")
def map_automatic_turn_detection(
self, value: OpenAIRealtimeTurnDetection
) -> AutomaticActivityDetection:
automatic_activity_dection = AutomaticActivityDetection()
if "create_response" in value and isinstance(value["create_response"], bool):
automatic_activity_dection["disabled"] = not value["create_response"]
else:
automatic_activity_dection["disabled"] = True
if "prefix_padding_ms" in value and isinstance(value["prefix_padding_ms"], int):
automatic_activity_dection["prefixPaddingMs"] = value["prefix_padding_ms"]
if "silence_duration_ms" in value and isinstance(
value["silence_duration_ms"], int
):
automatic_activity_dection["silenceDurationMs"] = value[
"silence_duration_ms"
]
return automatic_activity_dection
def map_openai_params(
self, optional_params: dict, non_default_params: dict
) -> dict:
if "generationConfig" not in optional_params:
optional_params["generationConfig"] = {}
for key, value in non_default_params.items():
if key == "instructions":
optional_params["systemInstruction"] = HttpxContentType(
role="user", parts=[{"text": value}]
)
elif key == "temperature":
optional_params["generationConfig"]["temperature"] = value
elif key == "max_response_output_tokens":
optional_params["generationConfig"]["maxOutputTokens"] = value
elif key == "modalities":
optional_params["generationConfig"]["responseModalities"] = [
modality.upper() for modality in cast(List[str], value)
]
elif key == "tools":
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
vertex_gemini_config = VertexGeminiConfig()
vertex_gemini_config._map_function(value)
optional_params["generationConfig"][
"tools"
] = vertex_gemini_config._map_function(value)
elif key == "input_audio_transcription" and value is not None:
optional_params["inputAudioTranscription"] = {}
elif key == "turn_detection":
value_typed = cast(OpenAIRealtimeTurnDetection, value)
transformed_audio_activity_config = self.map_automatic_turn_detection(
value_typed
)
if (
len(transformed_audio_activity_config) > 0
): # if the config is not empty, add it to the optional params
optional_params[
"realtimeInputConfig"
] = BidiGenerateContentRealtimeInputConfig(
automaticActivityDetection=transformed_audio_activity_config
)
if len(optional_params["generationConfig"]) == 0:
optional_params.pop("generationConfig")
return optional_params
def transform_realtime_request(
self,
message: str,
model: str,
session_configuration_request: Optional[str] = None,
) -> List[str]:
realtime_input_dict: BidiGenerateContentRealtimeInput = {}
try:
json_message = json.loads(message)
except json.JSONDecodeError:
if isinstance(message, bytes):
message_str = message.decode("utf-8", errors="replace")
else:
message_str = str(message)
raise ValueError(f"Invalid JSON message: {message_str}")
## HANDLE SESSION UPDATE ##
messages: List[str] = []
if "type" in json_message and json_message["type"] == "session.update":
client_session_configuration_request = self.map_openai_params(
optional_params={}, non_default_params=json_message["session"]
)
client_session_configuration_request["model"] = f"models/{model}"
messages.append(
json.dumps(
{
"setup": client_session_configuration_request,
}
)
)
# elif session_configuration_request is None:
# default_session_configuration_request = self.session_configuration_request(model)
# messages.append(default_session_configuration_request)
## HANDLE INPUT AUDIO BUFFER ##
if (
"type" in json_message
and json_message["type"] == "input_audio_buffer.append"
):
realtime_input_dict["audio"] = HttpxBlobType(
mimeType=self.get_audio_mime_type(), data=json_message["audio"]
)
else:
realtime_input_dict["text"] = message
if len(realtime_input_dict) != 1:
raise ValueError(
@ -82,9 +251,13 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
f" {list(realtime_input_dict.keys())}"
)
realtime_input_dict = encode_unserializable_types(realtime_input_dict)
realtime_input_dict = cast(
BidiGenerateContentRealtimeInput,
encode_unserializable_types(cast(Dict[str, object], realtime_input_dict)),
)
return json.dumps({"realtime_input": realtime_input_dict})
messages.append(json.dumps({"realtime_input": realtime_input_dict}))
return messages
def transform_session_created_event(
self,
@ -92,16 +265,21 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
logging_session_id: str,
session_configuration_request: Optional[str] = None,
) -> OpenAIRealtimeStreamSessionEvents:
if session_configuration_request is None:
raise ValueError(
"session_configuration_request is required for Gemini API calls"
)
if session_configuration_request:
session_configuration_request_dict: BidiGenerateContentSetup = json.loads(
session_configuration_request
).get("setup", {})
else:
session_configuration_request_dict = {}
session_configuration_request_dict = json.loads(session_configuration_request)
_model = session_configuration_request_dict.get("model") or model
_modalities = session_configuration_request_dict.get(
"generationConfig", {}
).get("responseModalities", ["TEXT"])
generation_config = (
session_configuration_request_dict.get("generationConfig", {}) or {}
)
gemini_modalities = generation_config.get("responseModalities", ["TEXT"])
_modalities = [
modality.lower() for modality in cast(List[str], gemini_modalities)
]
_system_instruction = session_configuration_request_dict.get(
"systemInstruction"
)
@ -112,7 +290,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
if _system_instruction is not None and isinstance(_system_instruction, str):
session["instructions"] = _system_instruction
if _model is not None and isinstance(_model, str):
session["model"] = _model
session["model"] = _model.strip(
"models/"
) # keep it consistent with how openai returns the model name
return OpenAIRealtimeStreamSessionEvents(
type="session.created",
@ -137,6 +317,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
response_id: str,
output_item_id: str,
conversation_id: str,
delta_type: ALL_DELTA_TYPES,
session_configuration_request: Optional[str] = None,
) -> List[OpenAIRealtimeEvents]:
if session_configuration_request is None:
@ -144,16 +325,19 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
"session_configuration_request is required for Gemini API calls"
)
session_configuration_request_dict = json.loads(session_configuration_request)
_modalities = session_configuration_request_dict.get(
session_configuration_request_dict: BidiGenerateContentSetup = json.loads(
session_configuration_request
).get("setup", {})
generation_config = session_configuration_request_dict.get(
"generationConfig", {}
).get("responseModalities", ["TEXT"])
_temperature = session_configuration_request_dict.get(
"generationConfig", {}
).get("temperature")
_max_output_tokens = session_configuration_request_dict.get(
"generationConfig", {}
).get("maxOutputTokens")
)
gemini_modalities = generation_config.get("responseModalities", ["TEXT"])
_modalities = [
modality.lower() for modality in cast(List[str], gemini_modalities)
]
_temperature = generation_config.get("temperature")
_max_output_tokens = generation_config.get("maxOutputTokens")
response_items: List[OpenAIRealtimeEvents] = []
@ -213,6 +397,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
part={
"type": "text",
"text": "",
}
if delta_type == "text"
else {
"type": "audio",
"transcript": "",
},
response_id=response_id,
)
@ -224,20 +413,25 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
message: BidiGenerateContentServerContent,
output_item_id: str,
response_id: str,
) -> OpenAIRealtimeResponseTextDelta:
delta_type: ALL_DELTA_TYPES,
) -> OpenAIRealtimeResponseDelta:
delta = ""
try:
if "modelTurn" in message and "parts" in message["modelTurn"]:
for part in message["modelTurn"]["parts"]:
if "text" in part:
delta += part["text"]
elif "inlineData" in part:
delta += part["inlineData"]["data"]
except Exception as e:
raise ValueError(
f"Error transforming content delta events: {e}, got message: {message}"
)
return OpenAIRealtimeResponseTextDelta(
type="response.text.delta",
return OpenAIRealtimeResponseDelta(
type="response.text.delta"
if delta_type == "text"
else "response.audio.delta",
content_index=0,
event_id="event_{}".format(uuid.uuid4()),
item_id=output_item_id,
@ -248,10 +442,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
def transform_content_done_event(
self,
delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]],
delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]],
current_output_item_id: Optional[str],
current_response_id: Optional[str],
) -> OpenAIRealtimeResponseTextDone:
delta_type: ALL_DELTA_TYPES,
) -> Union[OpenAIRealtimeResponseTextDone, OpenAIRealtimeResponseAudioDone]:
if delta_chunks:
delta = "".join([delta_chunk["delta"] for delta_chunk in delta_chunks])
else:
@ -260,21 +455,34 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
raise ValueError(
"current_output_item_id and current_response_id cannot be None for a 'done' event."
)
return OpenAIRealtimeResponseTextDone(
type="response.text.done",
content_index=0,
event_id="event_{}".format(uuid.uuid4()),
item_id=current_output_item_id,
output_index=0,
response_id=current_response_id,
text=delta,
)
if delta_type == "text":
return OpenAIRealtimeResponseTextDone(
type="response.text.done",
content_index=0,
event_id="event_{}".format(uuid.uuid4()),
item_id=current_output_item_id,
output_index=0,
response_id=current_response_id,
text=delta,
)
elif delta_type == "audio":
return OpenAIRealtimeResponseAudioDone(
type="response.audio.done",
content_index=0,
event_id="event_{}".format(uuid.uuid4()),
item_id=current_output_item_id,
output_index=0,
response_id=current_response_id,
)
def return_additional_content_done_events(
self,
current_output_item_id: Optional[str],
current_response_id: Optional[str],
delta_done_event: OpenAIRealtimeResponseTextDone,
delta_done_event: Union[
OpenAIRealtimeResponseTextDone, OpenAIRealtimeResponseAudioDone
],
delta_type: ALL_DELTA_TYPES,
) -> List[OpenAIRealtimeEvents]:
"""
- return response.content_part.done
@ -285,6 +493,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
"current_output_item_id and current_response_id cannot be None for a 'done' event."
)
returned_items: List[OpenAIRealtimeEvents] = []
delta_done_event_text = cast(Optional[str], delta_done_event.get("text"))
# response.content_part.done
response_content_part_done = OpenAIRealtimeContentPartDone(
type="response.content_part.done",
@ -292,9 +502,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
event_id="event_{}".format(uuid.uuid4()),
item_id=current_output_item_id,
output_index=0,
part={
"type": "text",
"text": delta_done_event["text"],
part={"type": "text", "text": delta_done_event_text}
if delta_done_event_text and delta_type == "text"
else {
"type": "audio",
"transcript": "", # gemini doesn't return transcript for audio
},
response_id=current_response_id,
)
@ -312,9 +524,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
"status": "completed",
"role": "assistant",
"content": [
{
"type": "text",
"text": delta_done_event["text"],
{"type": "text", "text": delta_done_event_text}
if delta_done_event_text and delta_type == "text"
else {
"type": "audio",
"transcript": "",
}
],
},
@ -336,8 +550,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
def update_current_delta_chunks(
self,
transformed_message: Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]],
current_delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]],
) -> Optional[List[OpenAIRealtimeResponseTextDelta]]:
current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]],
) -> Optional[List[OpenAIRealtimeResponseDelta]]:
try:
if isinstance(transformed_message, list):
current_delta_chunks = []
@ -345,7 +559,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
for event in transformed_message:
if event["type"] == "response.text.delta":
current_delta_chunks.append(
cast(OpenAIRealtimeResponseTextDelta, event)
cast(OpenAIRealtimeResponseDelta, event)
)
any_delta_chunk = True
if not any_delta_chunk:
@ -353,11 +567,13 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
None # reset current_delta_chunks if no delta chunks
)
else:
if transformed_message["type"] == "response.text.delta":
if (
transformed_message["type"] == "response.text.delta"
): # ONLY ACCUMULATE TEXT DELTA CHUNKS - AUDIO WILL CAUSE SERVER MEMORY ISSUES
if current_delta_chunks is None:
current_delta_chunks = []
current_delta_chunks.append(
cast(OpenAIRealtimeResponseTextDelta, transformed_message)
cast(OpenAIRealtimeResponseDelta, transformed_message)
)
else:
current_delta_chunks = None
@ -406,40 +622,41 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
message: BidiGenerateContentServerMessage,
current_response_id: Optional[str],
current_conversation_id: Optional[str],
current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]],
output_items: Optional[List[OpenAIRealtimeOutputItemDone]],
session_configuration_request: Optional[str] = None,
) -> OpenAIRealtimeDoneEvent:
if (
current_conversation_id is None
or current_response_id is None
or current_item_chunks is None
):
if current_conversation_id is None or current_response_id is None:
raise ValueError(
"current_conversation_id and current_response_id and current_item_chunks cannot be None for a 'done' event."
)
if session_configuration_request is None:
raise ValueError(
"session_configuration_request is required for Gemini API calls"
f"current_conversation_id and current_response_id must all be set for a 'done' event. Got=current_conversation_id: {current_conversation_id}, current_response_id: {current_response_id}"
)
session_configuration_request_dict = json.loads(session_configuration_request)
temperature = session_configuration_request_dict.get(
if session_configuration_request:
session_configuration_request_dict: BidiGenerateContentSetup = json.loads(
session_configuration_request
).get("setup", {})
else:
session_configuration_request_dict = {}
generation_config = session_configuration_request_dict.get(
"generationConfig", {}
).get("temperature")
max_output_tokens = session_configuration_request_dict.get(
"generationConfig", {}
).get("maxOutputTokens")
_modalities = session_configuration_request_dict.get(
"generationConfig", {}
).get("responseModalities", ["TEXT"])
_chat_completion_usage = VertexGeminiConfig()._calculate_usage(
completion_response=message,
)
temperature = generation_config.get("temperature")
max_output_tokens = generation_config.get("max_output_tokens")
gemini_modalities = generation_config.get("responseModalities", ["TEXT"])
_modalities = [
modality.lower() for modality in cast(List[str], gemini_modalities)
]
if "usageMetadata" in message:
_chat_completion_usage = VertexGeminiConfig()._calculate_usage(
completion_response=message,
)
else:
_chat_completion_usage = get_empty_usage()
responses_api_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage(
_chat_completion_usage,
)
return OpenAIRealtimeDoneEvent(
response_done_event = OpenAIRealtimeDoneEvent(
type="response.done",
event_id="event_{}".format(uuid.uuid4()),
response=OpenAIRealtimeResponseDoneObject(
@ -451,11 +668,121 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
else [],
conversation_id=current_conversation_id,
modalities=_modalities,
temperature=temperature,
max_output_tokens=max_output_tokens,
usage=responses_api_usage.model_dump(),
),
)
if temperature is not None:
response_done_event["response"]["temperature"] = temperature
if max_output_tokens is not None:
response_done_event["response"]["max_output_tokens"] = max_output_tokens
return response_done_event
def handle_openai_modality_event(
self,
openai_event: OpenAIRealtimeEventTypes,
json_message: dict,
realtime_response_transform_input: RealtimeResponseTransformInput,
delta_type: ALL_DELTA_TYPES,
) -> RealtimeModalityResponseTransformOutput:
current_output_item_id = realtime_response_transform_input[
"current_output_item_id"
]
current_response_id = realtime_response_transform_input["current_response_id"]
current_conversation_id = realtime_response_transform_input[
"current_conversation_id"
]
current_delta_chunks = realtime_response_transform_input["current_delta_chunks"]
session_configuration_request = realtime_response_transform_input[
"session_configuration_request"
]
returned_message: List[OpenAIRealtimeEvents] = []
if (
openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA
or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA
):
current_response_id = current_response_id or "resp_{}".format(uuid.uuid4())
if not current_output_item_id:
# send the list of standard 'new' content.delta events
current_output_item_id = "item_{}".format(uuid.uuid4())
current_conversation_id = current_conversation_id or "conv_{}".format(
uuid.uuid4()
)
returned_message = self.return_new_content_delta_events(
session_configuration_request=session_configuration_request,
response_id=current_response_id,
output_item_id=current_output_item_id,
conversation_id=current_conversation_id,
delta_type=delta_type,
)
# send the list of standard 'new' content.delta events
transformed_message = self.transform_content_delta_events(
BidiGenerateContentServerContent(**json_message["serverContent"]),
current_output_item_id,
current_response_id,
delta_type=delta_type,
)
returned_message.append(transformed_message)
elif (
openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE
or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE
):
transformed_content_done_event = self.transform_content_done_event(
current_output_item_id=current_output_item_id,
current_response_id=current_response_id,
delta_chunks=current_delta_chunks,
delta_type=delta_type,
)
returned_message = [transformed_content_done_event]
additional_items = self.return_additional_content_done_events(
current_output_item_id=current_output_item_id,
current_response_id=current_response_id,
delta_done_event=transformed_content_done_event,
delta_type=delta_type,
)
returned_message.extend(additional_items)
return {
"returned_message": returned_message,
"current_output_item_id": current_output_item_id,
"current_response_id": current_response_id,
"current_conversation_id": current_conversation_id,
"current_delta_chunks": current_delta_chunks,
"current_delta_type": delta_type,
}
def map_openai_event(
self,
key: str,
value: dict,
current_delta_type: Optional[ALL_DELTA_TYPES],
json_message: dict,
) -> OpenAIRealtimeEventTypes:
model_turn_event = value.get("modelTurn")
generation_complete_event = value.get("generationComplete")
openai_event: Optional[OpenAIRealtimeEventTypes] = None
if model_turn_event: # check if model turn event
openai_event = self.map_model_turn_event(model_turn_event)
elif generation_complete_event:
openai_event = self.map_generation_complete_event(
delta_type=current_delta_type
)
else:
# Check if this key or any nested key matches our mapping
for map_key, openai_event in MAP_GEMINI_FIELD_TO_OPENAI_EVENT.items():
if map_key == key or (
"." in map_key
and GeminiRealtimeConfig.get_nested_value(json_message, map_key)
is not None
):
openai_event = openai_event
break
if openai_event is None:
raise ValueError(f"Unknown openai event: {key}, value: {value}")
return openai_event
def transform_realtime_response(
self,
@ -490,91 +817,58 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
"session_configuration_request"
]
current_item_chunks = realtime_response_transform_input["current_item_chunks"]
returned_message: Optional[
Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]]
] = None
current_delta_type: Optional[
ALL_DELTA_TYPES
] = realtime_response_transform_input["current_delta_type"]
returned_message: List[OpenAIRealtimeEvents] = []
for key, value in json_message.items():
# Check if this key or any nested key matches our mapping
for map_key, openai_event in MAP_GEMINI_FIELD_TO_OPENAI_EVENT.items():
if map_key == key or (
"." in map_key
and GeminiRealtimeConfig.get_nested_value(json_message, map_key)
is not None
):
if openai_event == "session.created":
transformed_message = self.transform_session_created_event(
model,
logging_session_id,
realtime_response_transform_input[
"session_configuration_request"
],
)
returned_message = transformed_message
openai_event = self.map_openai_event(
key=key,
value=value,
current_delta_type=current_delta_type,
json_message=json_message,
)
elif openai_event == "response.text.delta":
# check if this is a new content.delta or a continuation of a previous content.delta
if not current_output_item_id:
# send the list of standard 'new' content.delta events
current_response_id = (
current_response_id or "resp_{}".format(uuid.uuid4())
)
current_output_item_id = "item_{}".format(uuid.uuid4())
current_conversation_id = (
current_conversation_id
or "conv_{}".format(uuid.uuid4())
)
response_items = self.return_new_content_delta_events(
session_configuration_request=session_configuration_request,
response_id=current_response_id,
output_item_id=current_output_item_id,
conversation_id=current_conversation_id,
)
transformed_message = self.transform_content_delta_events(
BidiGenerateContentServerContent(**json_message[key]), # type: ignore
current_output_item_id,
current_response_id,
)
response_items.append(transformed_message)
returned_message = response_items
else:
current_response_id = (
current_response_id or "resp_{}".format(uuid.uuid4())
)
# send the list of standard 'new' content.delta events
transformed_message = self.transform_content_delta_events(
BidiGenerateContentServerContent(**json_message[key]), # type: ignore
current_output_item_id,
current_response_id,
)
returned_message = transformed_message
elif openai_event == "response.text.done":
transformed_content_done_event = (
self.transform_content_done_event(
current_output_item_id=current_output_item_id,
current_response_id=current_response_id,
delta_chunks=current_delta_chunks,
)
)
returned_message = [transformed_content_done_event]
additional_items = self.return_additional_content_done_events(
current_output_item_id=current_output_item_id,
current_response_id=current_response_id,
delta_done_event=transformed_content_done_event,
)
returned_message.extend(additional_items)
elif openai_event == "response.done":
transformed_response_done_event = self.transform_response_done_event(
message=BidiGenerateContentServerMessage(**json_message), # type: ignore
current_response_id=current_response_id,
current_conversation_id=current_conversation_id,
session_configuration_request=session_configuration_request,
output_items=current_item_chunks,
)
returned_message = transformed_response_done_event
if returned_message is None:
if openai_event == OpenAIRealtimeEventTypes.SESSION_CREATED:
transformed_message = self.transform_session_created_event(
model,
logging_session_id,
realtime_response_transform_input["session_configuration_request"],
)
session_configuration_request = json.dumps(transformed_message)
returned_message.append(transformed_message)
elif openai_event == OpenAIRealtimeEventTypes.RESPONSE_DONE:
transformed_response_done_event = self.transform_response_done_event(
message=BidiGenerateContentServerMessage(**json_message), # type: ignore
current_response_id=current_response_id,
current_conversation_id=current_conversation_id,
session_configuration_request=session_configuration_request,
output_items=None,
)
returned_message.append(transformed_response_done_event)
elif (
openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA
or openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE
or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA
or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE
):
_returned_message = self.handle_openai_modality_event(
openai_event,
json_message,
realtime_response_transform_input,
delta_type="text" if "text" in openai_event.value else "audio",
)
returned_message.extend(_returned_message["returned_message"])
current_output_item_id = _returned_message["current_output_item_id"]
current_response_id = _returned_message["current_response_id"]
current_conversation_id = _returned_message["current_conversation_id"]
current_delta_chunks = _returned_message["current_delta_chunks"]
current_delta_type = _returned_message["current_delta_type"]
else:
raise ValueError(f"Unknown openai event: {openai_event}")
if len(returned_message) == 0:
if isinstance(message, bytes):
message_str = message.decode("utf-8", errors="replace")
else:
@ -596,12 +890,14 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
"current_delta_chunks": current_delta_chunks,
"current_conversation_id": current_conversation_id,
"current_item_chunks": current_item_chunks,
"current_delta_type": current_delta_type,
"session_configuration_request": session_configuration_request,
}
def requires_session_configuration(self) -> bool:
return True
def session_configuration_request(self, model: str) -> Optional[str]:
def session_configuration_request(self, model: str) -> str:
"""
```
@ -624,11 +920,20 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
}
```
"""
response_modalities: List[GeminiResponseModalities] = ["AUDIO"]
output_audio_transcription = False
# if "audio" in model: ## UNCOMMENT THIS WHEN AUDIO IS SUPPORTED
# output_audio_transcription = True
setup_config: BidiGenerateContentSetup = {
"model": f"models/{model}",
"generationConfig": {"responseModalities": response_modalities},
}
if output_audio_transcription:
setup_config["outputAudioTranscription"] = {}
return json.dumps(
{
"setup": {
"model": f"models/{model}",
"generationConfig": {"responseModalities": ["TEXT"]},
}
"setup": setup_config,
}
)

View file

@ -63,6 +63,7 @@ from litellm.types.llms.vertex_ai import (
from litellm.types.utils import (
ChatCompletionTokenLogprob,
ChoiceLogprobs,
CompletionTokensDetailsWrapper,
GenericStreamingChunk,
PromptTokensDetailsWrapper,
TopLogprob,
@ -803,10 +804,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
text_tokens: Optional[int] = None
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
reasoning_tokens: Optional[int] = None
response_tokens: Optional[int] = None
response_tokens_details: Optional[CompletionTokensDetailsWrapper] = None
if "cachedContentTokenCount" in completion_response["usageMetadata"]:
cached_tokens = completion_response["usageMetadata"][
"cachedContentTokenCount"
]
## GEMINI LIVE API ONLY PARAMS ##
if "responseTokenCount" in completion_response["usageMetadata"]:
response_tokens = completion_response["usageMetadata"]["responseTokenCount"]
if "responseTokensDetails" in completion_response["usageMetadata"]:
response_tokens_details = CompletionTokensDetailsWrapper()
for detail in completion_response["usageMetadata"]["responseTokensDetails"]:
if detail["modality"] == "TEXT":
response_tokens_details.text_tokens = detail["tokenCount"]
elif detail["modality"] == "AUDIO":
response_tokens_details.audio_tokens = detail["tokenCount"]
#########################################################
if "promptTokensDetails" in completion_response["usageMetadata"]:
for detail in completion_response["usageMetadata"]["promptTokensDetails"]:
if detail["modality"] == "AUDIO":
@ -823,7 +839,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
text_tokens=text_tokens,
)
completion_tokens = completion_response["usageMetadata"].get(
completion_tokens = response_tokens or completion_response["usageMetadata"].get(
"candidatesTokenCount", 0
)
if (
@ -842,6 +858,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
total_tokens=completion_response["usageMetadata"].get("totalTokenCount", 0),
prompt_tokens_details=prompt_tokens_details,
reasoning_tokens=reasoning_tokens,
completion_tokens_details=response_tokens_details,
)
return usage
@ -1637,7 +1654,7 @@ class ModelResponseIterator:
"reasoning_tokens": processed_chunk["usageMetadata"].get(
"thoughtsTokenCount", 0
)
}
},
)
returned_chunk = GenericStreamingChunk(

File diff suppressed because one or more lines are too long

View file

@ -6,6 +6,10 @@ model_list:
litellm_params:
model: gpt-4o-mini
api_key: os.environ/OPENAI_API_KEY
- model_name: "gpt-4o-realtime-preview"
litellm_params:
model: gpt-4o-realtime-preview-2024-10-01
api_key: os.environ/OPENAI_API_KEY
- model_name: "bedrock-nova"
litellm_params:
model: us.amazon.nova-pro-v1:0

View file

@ -3,7 +3,13 @@ from typing import Any, Dict, Iterable, List, Literal, Optional, Union
from typing_extensions import Required, TypedDict
from .vertex_ai import HttpxContentType, UsageMetadata
from .vertex_ai import (
GenerationConfig,
HttpxBlobType,
HttpxContentType,
Tools,
UsageMetadata,
)
class GeminiFilesState(Enum):
@ -72,3 +78,75 @@ class BidiGenerateContentServerMessage(TypedDict, total=False):
setupComplete: dict
"""Output only. The setup complete message."""
class BidiGenerateContentRealtimeInput(TypedDict, total=False):
text: str
"""The text to be sent to the model."""
audio: HttpxBlobType
"""The audio to be sent to the model."""
video: HttpxBlobType
"""The video to be sent to the model."""
audioStreamEnd: bool
"""Output only. If true, indicates that the audio stream has ended."""
activityStart: bool
"""Output only. If true, indicates that the activity has started."""
activityEnd: bool
"""Output only. If true, indicates that the activity has ended."""
StartOfSpeechSensitivityEnum = Literal[
"START_SENSITIVITY_UNSPECIFIED", "START_SENSITIVITY_HIGH", "START_SENSITIVITY_LOW"
]
EndOfSpeechSensitivityEnum = Literal[
"END_SENSITIVITY_UNSPECIFIED", "END_SENSITIVITY_HIGH", "END_SENSITIVITY_LOW"
]
class AutomaticActivityDetection(TypedDict, total=False):
disabled: bool
startOfSpeechSensitivity: StartOfSpeechSensitivityEnum
prefixPaddingMs: int
endOfSpeechSensitivity: EndOfSpeechSensitivityEnum
silenceDurationMs: int
class BidiGenerateContentRealtimeInputConfig(TypedDict, total=False):
automaticActivityDetection: AutomaticActivityDetection
class BidiGenerateContentSetup(TypedDict, total=False):
model: str
"""The model to be used for the realtime session."""
generationConfig: GenerationConfig
"""The generation config to be used for the realtime session."""
systemInstruction: HttpxContentType
"""The system instruction to be used for the realtime session."""
tools: List[Tools]
"""The tools to be used for the realtime session."""
realtimeInputConfig: dict
"""The realtime config to be used for the realtime session."""
sessionResumption: dict
"""The session resumption to be used for the realtime session."""
sessionResumptionConfig: dict
"""The session resumption config to be used for the realtime session."""
contextWindowCompression: dict
"""The context window compression to be used for the realtime session."""
inputAudioTranscription: dict
"""The input audio transcription to be used for the realtime session."""
outputAudioTranscription: dict
"""The output audio transcription to be used for the realtime session."""

View file

@ -1353,7 +1353,7 @@ class OpenAIRealtimeStreamResponseOutputItemContent(TypedDict, total=False):
"""The text content, used for 'input_text' and 'text' content types"""
transcript: str
"""The transcript content, used for 'input_audio' content types"""
type: Literal["input_audio", "input_text", "text", "item_reference"]
type: Literal["input_audio", "input_text", "text", "item_reference", "audio"]
"""The type of content"""
@ -1443,14 +1443,14 @@ class OpenAIRealtimeResponseContentPartAdded(TypedDict):
response_id: str
class OpenAIRealtimeResponseTextDelta(TypedDict):
class OpenAIRealtimeResponseDelta(TypedDict):
content_index: int
delta: str
event_id: str
item_id: str
output_index: int
response_id: str
type: Literal["response.text.delta"]
type: Union[Literal["response.text.delta"], Literal["response.audio.delta"]]
class OpenAIRealtimeResponseTextDone(TypedDict):
@ -1463,6 +1463,15 @@ class OpenAIRealtimeResponseTextDone(TypedDict):
type: Literal["response.text.done"]
class OpenAIRealtimeResponseAudioDone(TypedDict):
content_index: int
event_id: str
item_id: str
output_index: int
response_id: str
type: Literal["response.audio.done"]
class OpenAIRealtimeContentPartDone(TypedDict):
content_index: int
event_id: str
@ -1503,6 +1512,17 @@ class OpenAIRealtimeDoneEvent(TypedDict):
type: Literal["response.done"]
class OpenAIRealtimeEventTypes(Enum):
SESSION_CREATED = "session.created"
RESPONSE_TEXT_DELTA = "response.text.delta"
RESPONSE_AUDIO_DELTA = "response.audio.delta"
RESPONSE_TEXT_DONE = "response.text.done"
RESPONSE_AUDIO_DONE = "response.audio.done"
RESPONSE_DONE = "response.done"
RESPONSE_OUTPUT_ITEM_ADDED = "response.output_item.added"
RESPONSE_CONTENT_PART_ADDED = "response.content_part.added"
OpenAIRealtimeEvents = Union[
OpenAIRealtimeStreamResponseBaseObject,
OpenAIRealtimeStreamSessionEvents,
@ -1510,8 +1530,9 @@ OpenAIRealtimeEvents = Union[
OpenAIRealtimeResponseContentPartAdded,
OpenAIRealtimeConversationItemCreated,
OpenAIRealtimeConversationCreated,
OpenAIRealtimeResponseTextDelta,
OpenAIRealtimeResponseDelta,
OpenAIRealtimeResponseTextDone,
OpenAIRealtimeResponseAudioDone,
OpenAIRealtimeContentPartDone,
OpenAIRealtimeOutputItemDone,
OpenAIRealtimeDoneEvent,
@ -1610,3 +1631,13 @@ class OpenAIWebSearchUserLocation(TypedDict):
class OpenAIWebSearchOptions(TypedDict, total=False):
search_context_size: Optional[Literal["low", "medium", "high"]]
user_location: Optional[OpenAIWebSearchUserLocation]
class OpenAIRealtimeTurnDetection(TypedDict, total=False):
create_response: bool
eagerness: str
interrupt_response: bool
prefix_padding_ms: int
silence_duration_ms: int
threshold: int
type: str

View file

@ -173,6 +173,9 @@ class GeminiThinkingConfig(TypedDict, total=False):
thinkingBudget: int
GeminiResponseModalities = Literal["TEXT", "IMAGE", "AUDIO", "VIDEO"]
class GenerationConfig(TypedDict, total=False):
temperature: float
top_p: float
@ -187,7 +190,7 @@ class GenerationConfig(TypedDict, total=False):
seed: int
responseLogprobs: bool
logprobs: int
responseModalities: List[Literal["TEXT", "IMAGE", "AUDIO", "VIDEO"]]
responseModalities: List[GeminiResponseModalities]
thinkingConfig: GeminiThinkingConfig
@ -218,9 +221,11 @@ class UsageMetadata(TypedDict, total=False):
promptTokenCount: int
totalTokenCount: int
candidatesTokenCount: int
responseTokenCount: int
cachedContentTokenCount: int
promptTokensDetails: List[PromptTokensDetails]
thoughtsTokenCount: int
responseTokensDetails: List[PromptTokensDetails]
class CachedContent(TypedDict, total=False):

View file

@ -1,11 +1,13 @@
from typing import List, Optional, TypedDict, Union
from typing import List, Literal, Optional, TypedDict, Union
from .llms.openai import (
OpenAIRealtimeEvents,
OpenAIRealtimeOutputItemDone,
OpenAIRealtimeResponseTextDelta,
OpenAIRealtimeResponseDelta,
)
ALL_DELTA_TYPES = Literal["text", "audio"]
class RealtimeResponseTransformInput(TypedDict):
session_configuration_request: Optional[str]
@ -15,15 +17,27 @@ class RealtimeResponseTransformInput(TypedDict):
current_response_id: Optional[
str
] # used to check if this is a new content.delta or a continuation of a previous content.delta
current_delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]]
current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]]
current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]]
current_conversation_id: Optional[str]
current_delta_type: Optional[ALL_DELTA_TYPES]
class RealtimeResponseTypedDict(TypedDict):
response: Union[OpenAIRealtimeEvents, List[OpenAIRealtimeEvents]]
current_output_item_id: Optional[str]
current_response_id: Optional[str]
current_delta_chunks: Optional[List[OpenAIRealtimeResponseTextDelta]]
current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]]
current_conversation_id: Optional[str]
current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]]
current_delta_type: Optional[ALL_DELTA_TYPES]
session_configuration_request: Optional[str]
class RealtimeModalityResponseTransformOutput(TypedDict):
returned_message: List[OpenAIRealtimeEvents]
current_output_item_id: Optional[str]
current_response_id: Optional[str]
current_conversation_id: Optional[str]
current_delta_chunks: Optional[List[OpenAIRealtimeResponseDelta]]
current_delta_type: Optional[ALL_DELTA_TYPES]

View file

@ -6876,3 +6876,11 @@ def jsonify_tools(tools: List[Any]) -> List[Dict]:
if isinstance(tool, dict):
new_tools.append(tool)
return new_tools
def get_empty_usage() -> Usage:
return Usage(
prompt_tokens=0,
completion_tokens=0,
total_tokens=0,
)

View file

@ -40,9 +40,12 @@ def test_gemini_realtime_transformation_session_created():
"current_conversation_id": None,
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": None,
},
)
assert transformed_message["response"]["type"] == "session.created"
print(transformed_message)
assert transformed_message["response"][0]["type"] == "session.created"
def test_gemini_realtime_transformation_content_delta():
@ -80,6 +83,7 @@ def test_gemini_realtime_transformation_content_delta():
"current_conversation_id": None,
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": None,
},
)
transformed_message = returned_object["response"]
@ -105,3 +109,121 @@ def test_gemini_realtime_transformation_content_delta():
event["item_id"] for event in transformed_message if "item_id" in event
]
assert len(set(output_item_ids)) == 1
def test_gemini_model_turn_event_mapping():
from litellm.types.llms.openai import OpenAIRealtimeEventTypes
config = GeminiRealtimeConfig()
assert config is not None
model_turn_event = {"parts": [{"text": "Hello, world!"}]}
openai_event = config.map_model_turn_event(model_turn_event)
assert openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA
model_turn_event = {
"parts": [{"inlineData": {"mimeType": "audio/pcm", "data": "..."}}]
}
openai_event = config.map_model_turn_event(model_turn_event)
assert openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA
model_turn_event = {
"parts": [
{
"text": "Hello, world!",
"inlineData": {"mimeType": "audio/pcm", "data": "..."},
}
]
}
openai_event = config.map_model_turn_event(model_turn_event)
assert openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA
def test_gemini_realtime_transformation_audio_delta():
from litellm.types.llms.openai import OpenAIRealtimeEventTypes
config = GeminiRealtimeConfig()
assert config is not None
session_configuration_request = {
"model": "gemini-1.5-flash",
"generationConfig": {"responseModalities": ["AUDIO"]},
}
session_configuration_request_str = json.dumps(session_configuration_request)
audio_delta_event = {
"serverContent": {
"modelTurn": {
"parts": [
{"inlineData": {"mimeType": "audio/pcm", "data": "my-audio-data"}}
]
}
}
}
result = config.transform_realtime_response(
json.dumps(audio_delta_event),
"gemini-1.5-flash",
MagicMock(),
realtime_response_transform_input={
"session_configuration_request": session_configuration_request_str,
"current_output_item_id": None,
"current_response_id": None,
"current_conversation_id": None,
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": None,
},
)
print(result)
responses = result["response"]
contains_audio_delta = False
for response in responses:
if response["type"] == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA.value:
contains_audio_delta = True
break
assert contains_audio_delta, "Expected audio delta event"
def test_gemini_realtime_transformation_generation_complete():
from litellm.types.llms.openai import OpenAIRealtimeEventTypes
config = GeminiRealtimeConfig()
assert config is not None
session_configuration_request = {
"model": "gemini-1.5-flash",
"generationConfig": {"responseModalities": ["AUDIO"]},
}
session_configuration_request_str = json.dumps(session_configuration_request)
audio_delta_event = {"serverContent": {"generationComplete": True}}
result = config.transform_realtime_response(
json.dumps(audio_delta_event),
"gemini-1.5-flash",
MagicMock(),
realtime_response_transform_input={
"session_configuration_request": session_configuration_request_str,
"current_output_item_id": "my-output-item-id",
"current_response_id": "my-response-id",
"current_conversation_id": None,
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": "audio",
},
)
print(result)
responses = result["response"]
contains_audio_done_event = False
for response in responses:
if response["type"] == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE.value:
contains_audio_delta = True
break
assert contains_audio_delta, "Expected audio delta event"

View file

@ -316,23 +316,21 @@ def test_vertex_ai_candidate_token_count_inclusive(
assert usage.completion_tokens == expected_usage.completion_tokens
assert usage.total_tokens == expected_usage.total_tokens
def test_streaming_chunk_includes_reasoning_tokens():
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ModelResponseIterator
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
# Simulate a streaming chunk as would be received from Gemini
chunk = {
"candidates": [
{
"content": {
"parts": [{"text": "Hello"}]
}
}
],
"candidates": [{"content": {"parts": [{"text": "Hello"}]}}],
"usageMetadata": {
"promptTokenCount": 5,
"candidatesTokenCount": 7,
"totalTokenCount": 12,
"thoughtsTokenCount": 3,
}
},
}
iterator = ModelResponseIterator(streaming_response=[], sync_stream=True)
streaming_chunk = iterator.chunk_parser(chunk)
@ -340,7 +338,10 @@ def test_streaming_chunk_includes_reasoning_tokens():
assert streaming_chunk["usage"]["prompt_tokens"] == 5
assert streaming_chunk["usage"]["completion_tokens"] == 7
assert streaming_chunk["usage"]["total_tokens"] == 12
assert streaming_chunk["usage"]["completion_tokens_details"]["reasoning_tokens"] == 3
assert (
streaming_chunk["usage"]["completion_tokens_details"]["reasoning_tokens"] == 3
)
def test_check_finish_reason():
config = VertexGeminiConfig()
@ -350,3 +351,27 @@ def test_check_finish_reason():
config._check_finish_reason(chat_completion_message=None, finish_reason=k)
== v
)
def test_vertex_ai_usage_metadata_response_token_count():
"""For Gemini Live API"""
from litellm.types.utils import PromptTokensDetailsWrapper
v = VertexGeminiConfig()
usage_metadata = {
"promptTokenCount": 57,
"responseTokenCount": 74,
"totalTokenCount": 131,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 57}],
"responseTokensDetails": [{"modality": "TEXT", "tokenCount": 74}],
}
usage_metadata = UsageMetadata(**usage_metadata)
result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata})
print("result", result)
assert result.prompt_tokens == 57
assert result.completion_tokens == 74
assert result.total_tokens == 131
assert result.prompt_tokens_details.text_tokens == 57
assert result.prompt_tokens_details.audio_tokens is None
assert result.prompt_tokens_details.cached_tokens is None
assert result.completion_tokens_details.text_tokens == 74

View file

@ -3161,3 +3161,5 @@ async def test_bedrock_max_completion_tokens(model: str):
"system": [],
"inferenceConfig": {"maxTokens": 10},
}