build: merge in https://github.com/BerriAI/litellm/pull/10909
Closes https://github.com/BerriAI/litellm/pull/10909
This commit is contained in:
parent
cc626ad3ec
commit
5c5699b65d
@ -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
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
}
|
||||
)
|
||||
|
||||
@ -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
@ -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
|
||||
|
||||
@ -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."""
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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]
|
||||
|
||||
@ -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,
|
||||
)
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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
|
||||
|
||||
@ -3161,3 +3161,5 @@ async def test_bedrock_max_completion_tokens(model: str):
|
||||
"system": [],
|
||||
"inferenceConfig": {"maxTokens": 10},
|
||||
}
|
||||
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user