mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
699 lines
27 KiB
Python
699 lines
27 KiB
Python
"""Abstraction function for OpenAI's realtime API"""
|
|
|
|
import asyncio
|
|
import os
|
|
from collections.abc import Mapping
|
|
from types import MappingProxyType
|
|
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
|
|
|
import litellm
|
|
from litellm.constants import (
|
|
AZURE_OPENAI_AUDIO_PROVIDERS,
|
|
REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
|
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
|
request_timeout,
|
|
)
|
|
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
|
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
|
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
|
from litellm.llms.xai.common_utils import XAIModelInfo
|
|
from litellm.secret_managers.main import get_secret_str
|
|
from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES, VertexAccessTokenResolver
|
|
from litellm.types.realtime import (
|
|
RealtimeClientSecretRequest,
|
|
RealtimeExpiresAfter,
|
|
RealtimeQueryParams,
|
|
RealtimeSessionConfig,
|
|
RealtimeTranscriptionSessionRequest,
|
|
)
|
|
from litellm.types.router import GenericLiteLLMParams
|
|
from litellm.types.utils import CallTypes, LlmProviders
|
|
from litellm.utils import ProviderConfigManager
|
|
|
|
from ..litellm_core_utils.get_litellm_params import get_litellm_params
|
|
from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
|
from ..llms.azure.common_utils import get_azure_ad_token
|
|
from ..llms.azure.realtime.handler import AzureOpenAIRealtime
|
|
from ..llms.bedrock.realtime.handler import BedrockRealtime
|
|
from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
|
|
from ..llms.openai.realtime.handler import OpenAIRealtime
|
|
from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
|
|
from ..llms.vertex_ai.vertex_llm_base import VertexBase
|
|
from ..llms.xai.realtime.handler import XAIRealtime
|
|
from ..utils import client as wrapper_client
|
|
|
|
if TYPE_CHECKING:
|
|
from fastapi import WebSocket
|
|
|
|
from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig
|
|
|
|
azure_realtime: Final = AzureOpenAIRealtime()
|
|
openai_realtime: Final = OpenAIRealtime()
|
|
bedrock_realtime: Final = BedrockRealtime()
|
|
xai_realtime: Final = XAIRealtime()
|
|
vertex_llm_base: Final = VertexBase()
|
|
base_llm_http_handler = BaseLLMHTTPHandler()
|
|
_EMPTY_MODEL_PARAMS: Final[Mapping[str, Any]] = MappingProxyType({})
|
|
|
|
|
|
def _with_resolved_session_model(session: dict[str, object], model_name: str) -> dict[str, object]:
|
|
if "model" not in session:
|
|
return session
|
|
return {**session, "model": model_name}
|
|
|
|
|
|
def _build_litellm_metadata(kwargs: dict) -> dict:
|
|
"""Build the litellm_metadata dict for guardrail checking (internal only, not forwarded to provider)."""
|
|
metadata: Final[dict] = {**(kwargs.get("litellm_metadata") or {})}
|
|
guardrails: Final = (kwargs.get("metadata") or {}).get("guardrails") or kwargs.get("guardrails") or []
|
|
if guardrails:
|
|
metadata["guardrails"] = guardrails
|
|
return metadata
|
|
|
|
|
|
def _get_realtime_http_provider_config(
|
|
custom_llm_provider: str,
|
|
dynamic_api_base: str | None,
|
|
dynamic_api_key: str | None,
|
|
litellm_params: GenericLiteLLMParams,
|
|
) -> tuple["BaseRealtimeHTTPConfig | None", str, str]:
|
|
"""
|
|
Return (provider_config, resolved_api_base, resolved_api_key) for the
|
|
realtime HTTP endpoints (client_secrets / realtime_calls).
|
|
|
|
Uses ProviderConfigManager so each provider keeps its credential-resolution
|
|
and URL-construction logic in its own transformation class.
|
|
"""
|
|
from litellm.llms.base_llm.realtime.http_transformation import (
|
|
BaseRealtimeHTTPConfig,
|
|
)
|
|
|
|
provider_config: BaseRealtimeHTTPConfig | None = None
|
|
if custom_llm_provider in LlmProviders._member_map_.values():
|
|
provider_config = ProviderConfigManager.get_provider_realtime_http_config(
|
|
model="",
|
|
provider=LlmProviders(custom_llm_provider),
|
|
)
|
|
|
|
raw_api_base: Final = dynamic_api_base or litellm_params.api_base
|
|
raw_api_key: Final = dynamic_api_key or litellm_params.api_key
|
|
|
|
if provider_config is not None:
|
|
resolved_api_base = provider_config.get_api_base(api_base=raw_api_base)
|
|
resolved_api_key = provider_config.get_api_key(api_key=raw_api_key)
|
|
else:
|
|
# Fallback for providers without a dedicated HTTP config (treated as OpenAI-compatible).
|
|
resolved_api_base = raw_api_base or litellm.api_base or "https://api.openai.com"
|
|
resolved_api_key = (
|
|
raw_api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") or ""
|
|
)
|
|
|
|
return provider_config, resolved_api_base.rstrip("/"), resolved_api_key
|
|
|
|
|
|
@wrapper_client
|
|
async def acreate_realtime_client_secret(
|
|
model: str | None = None,
|
|
session: dict[str, Any] | None = None,
|
|
expires_after: dict[str, Any] | None = None,
|
|
timeout: float | None = None,
|
|
**kwargs,
|
|
):
|
|
req: Final = RealtimeClientSecretRequest(
|
|
model=model,
|
|
session=RealtimeSessionConfig(**session) if session else None,
|
|
expires_after=RealtimeExpiresAfter(**expires_after) if expires_after else None,
|
|
)
|
|
model_name = (req.session.model if req.session is not None else None) or req.model or "gpt-4o-realtime-preview"
|
|
litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj")
|
|
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
|
|
|
(
|
|
model_name,
|
|
custom_llm_provider,
|
|
dynamic_api_key,
|
|
dynamic_api_base,
|
|
) = get_llm_provider(
|
|
model=model_name,
|
|
api_base=litellm_params.api_base,
|
|
api_key=litellm_params.api_key,
|
|
)
|
|
(
|
|
provider_config,
|
|
resolved_api_base,
|
|
resolved_api_key,
|
|
) = _get_realtime_http_provider_config(
|
|
custom_llm_provider=custom_llm_provider,
|
|
dynamic_api_base=dynamic_api_base,
|
|
dynamic_api_key=dynamic_api_key,
|
|
litellm_params=litellm_params,
|
|
)
|
|
litellm_logging_obj.update_from_kwargs(
|
|
kwargs=kwargs,
|
|
model=model_name,
|
|
optional_params={"expires_after": expires_after, "session": session},
|
|
litellm_params={"api_base": resolved_api_base},
|
|
custom_llm_provider=custom_llm_provider,
|
|
)
|
|
request_data: Final = req.model_dump(exclude_none=True, exclude={"model"})
|
|
if isinstance(request_data.get("session"), dict):
|
|
request_data["session"] = _with_resolved_session_model(request_data["session"], model_name)
|
|
return await base_llm_http_handler.async_realtime_client_secret_handler(
|
|
api_base=resolved_api_base,
|
|
api_key=resolved_api_key,
|
|
request_data=request_data,
|
|
logging_obj=litellm_logging_obj,
|
|
timeout=timeout or request_timeout,
|
|
provider_config=provider_config,
|
|
model=model_name,
|
|
extra_headers=kwargs.get("extra_headers"),
|
|
client=kwargs.get("client"),
|
|
api_version=litellm_params.api_version,
|
|
)
|
|
|
|
|
|
@wrapper_client
|
|
async def acreate_realtime_transcription_session(
|
|
model: str | None = None,
|
|
transcription_session: dict[str, Any] | None = None,
|
|
timeout: float | None = None,
|
|
**kwargs,
|
|
):
|
|
"""
|
|
Create an ephemeral transcription session via POST
|
|
/v1/realtime/transcription_sessions.
|
|
|
|
``transcription_session`` is the upstream request body (input_audio_format,
|
|
input_audio_transcription, turn_detection, …). ``model`` is a LiteLLM-only
|
|
routing hint; the provider model lives in
|
|
``transcription_session.input_audio_transcription.model``.
|
|
"""
|
|
req: Final = RealtimeTranscriptionSessionRequest(
|
|
model=model,
|
|
**(transcription_session or {}),
|
|
)
|
|
model_name = req.resolved_model() or "gpt-realtime-whisper"
|
|
litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj")
|
|
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
|
|
|
(
|
|
model_name,
|
|
custom_llm_provider,
|
|
dynamic_api_key,
|
|
dynamic_api_base,
|
|
) = get_llm_provider(
|
|
model=model_name,
|
|
api_base=litellm_params.api_base,
|
|
api_key=litellm_params.api_key,
|
|
)
|
|
(
|
|
provider_config,
|
|
resolved_api_base,
|
|
resolved_api_key,
|
|
) = _get_realtime_http_provider_config(
|
|
custom_llm_provider=custom_llm_provider,
|
|
dynamic_api_base=dynamic_api_base,
|
|
dynamic_api_key=dynamic_api_key,
|
|
litellm_params=litellm_params,
|
|
)
|
|
litellm_logging_obj.update_from_kwargs(
|
|
kwargs=kwargs,
|
|
model=model_name,
|
|
optional_params={"transcription_session": transcription_session},
|
|
litellm_params={"api_base": resolved_api_base},
|
|
custom_llm_provider=custom_llm_provider,
|
|
)
|
|
request_data: Final = req.model_dump(exclude_none=True, exclude={"model"})
|
|
# Ensure the upstream body's input_audio_transcription.model matches the
|
|
# authorized routing model. This prevents a caller from supplying an allowed
|
|
# top-level model for auth while sneaking a different model into the nested
|
|
# transcription config that gets forwarded to the provider.
|
|
if isinstance(request_data.get("input_audio_transcription"), dict):
|
|
request_data["input_audio_transcription"]["model"] = model_name
|
|
return await base_llm_http_handler.async_realtime_transcription_session_handler(
|
|
api_base=resolved_api_base,
|
|
api_key=resolved_api_key,
|
|
request_data=request_data,
|
|
logging_obj=litellm_logging_obj,
|
|
timeout=timeout or request_timeout,
|
|
provider_config=provider_config,
|
|
model=model_name,
|
|
extra_headers=kwargs.get("extra_headers"),
|
|
client=kwargs.get("client"),
|
|
api_version=litellm_params.api_version,
|
|
)
|
|
|
|
|
|
@wrapper_client
|
|
async def arealtime_calls(
|
|
openai_ephemeral_key: str,
|
|
sdp_body: bytes,
|
|
model: str | None = None,
|
|
session: dict[str, Any] | None = None,
|
|
timeout: float | None = None,
|
|
**kwargs,
|
|
):
|
|
model_name = model or "gpt-4o-realtime-preview"
|
|
litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj")
|
|
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
|
|
|
(
|
|
model_name,
|
|
custom_llm_provider,
|
|
dynamic_api_key,
|
|
dynamic_api_base,
|
|
) = get_llm_provider(
|
|
model=model_name,
|
|
api_base=litellm_params.api_base,
|
|
api_key=litellm_params.api_key,
|
|
)
|
|
provider_config, resolved_api_base, _ = _get_realtime_http_provider_config(
|
|
custom_llm_provider=custom_llm_provider,
|
|
dynamic_api_base=dynamic_api_base,
|
|
dynamic_api_key=dynamic_api_key,
|
|
litellm_params=litellm_params,
|
|
)
|
|
if session is not None:
|
|
session = _with_resolved_session_model(session, model_name)
|
|
litellm_logging_obj.update_from_kwargs(
|
|
kwargs=kwargs,
|
|
model=model_name,
|
|
optional_params={"realtime_calls": True, "session": session},
|
|
litellm_params={"api_base": resolved_api_base},
|
|
custom_llm_provider=custom_llm_provider,
|
|
)
|
|
return await base_llm_http_handler.async_realtime_calls_handler(
|
|
api_base=resolved_api_base,
|
|
openai_ephemeral_key=openai_ephemeral_key,
|
|
sdp_body=sdp_body,
|
|
logging_obj=litellm_logging_obj,
|
|
timeout=timeout or request_timeout,
|
|
provider_config=provider_config,
|
|
model=model_name,
|
|
session_config=session,
|
|
extra_headers=kwargs.get("extra_headers"),
|
|
client=kwargs.get("client"),
|
|
api_version=litellm_params.api_version,
|
|
)
|
|
|
|
|
|
async def vertex_access_token_resolver(
|
|
credentials: VERTEX_CREDENTIALS_TYPES | None,
|
|
project_id: str | None,
|
|
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
|
) -> tuple[str, str]:
|
|
return await vertex_llm_base._ensure_access_token_async(
|
|
credentials=credentials,
|
|
project_id=project_id,
|
|
custom_llm_provider=custom_llm_provider,
|
|
)
|
|
|
|
|
|
async def _resolve_vertex_access_token_bounded(
|
|
credentials: VERTEX_CREDENTIALS_TYPES | None,
|
|
project_id: str | None,
|
|
resolver: VertexAccessTokenResolver,
|
|
timeout_seconds: float,
|
|
) -> tuple[str, str]:
|
|
try:
|
|
return await asyncio.wait_for(
|
|
resolver(
|
|
credentials=credentials,
|
|
project_id=project_id,
|
|
custom_llm_provider="vertex_ai",
|
|
),
|
|
timeout=timeout_seconds,
|
|
)
|
|
except asyncio.TimeoutError as e:
|
|
raise ValueError(
|
|
"Vertex AI realtime: timed out fetching Google OAuth access token after "
|
|
f"{timeout_seconds}s; check network egress from the proxy "
|
|
"to the OAuth token endpoint (oauth2.googleapis.com)"
|
|
) from e
|
|
|
|
|
|
@wrapper_client
|
|
async def _arealtime(
|
|
model: str,
|
|
websocket: "WebSocket", # fastapi websocket
|
|
api_base: str | None = None,
|
|
api_key: str | None = None,
|
|
api_version: str | None = None,
|
|
azure_ad_token: str | None = None,
|
|
client: object | None = None,
|
|
timeout: float | None = None,
|
|
query_params: RealtimeQueryParams | None = None,
|
|
**kwargs,
|
|
):
|
|
"""
|
|
Private function to handle the realtime API call.
|
|
|
|
For PROXY use only.
|
|
"""
|
|
headers = cast(dict | None, kwargs.get("headers"))
|
|
extra_headers: Final = cast(dict | None, kwargs.get("extra_headers"))
|
|
if headers is None:
|
|
headers = {}
|
|
if extra_headers is not None:
|
|
headers.update(extra_headers)
|
|
litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj")
|
|
user: Final = kwargs.get("user", None)
|
|
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
|
|
|
litellm_params_dict: Final = {**get_litellm_params(**kwargs), CallTypes.arealtime.value: True}
|
|
|
|
model, _custom_llm_provider, dynamic_api_key, dynamic_api_base = get_llm_provider(
|
|
model=model,
|
|
api_base=api_base,
|
|
api_key=api_key,
|
|
)
|
|
|
|
# If the client supplied `model` in the URL, ensure it uses the normalized
|
|
# provider model (no proxy aliases). If they omitted it, preserve that shape
|
|
# for transcription-only sessions like OpenAI's `?intent=transcription`.
|
|
if query_params is not None:
|
|
query_params = {**query_params}
|
|
if "model" in query_params:
|
|
query_params["model"] = model
|
|
|
|
litellm_logging_obj.update_from_kwargs(
|
|
kwargs=kwargs,
|
|
model=model,
|
|
user=user,
|
|
optional_params={},
|
|
litellm_params=litellm_params_dict,
|
|
custom_llm_provider=_custom_llm_provider,
|
|
)
|
|
|
|
provider_config: BaseRealtimeConfig | None = None
|
|
if _custom_llm_provider in LlmProviders._member_map_.values():
|
|
provider_config = ProviderConfigManager.get_provider_realtime_config(
|
|
model=model,
|
|
provider=LlmProviders(_custom_llm_provider),
|
|
)
|
|
if provider_config is not None:
|
|
await base_llm_http_handler.async_realtime(
|
|
model=model,
|
|
websocket=websocket,
|
|
logging_obj=litellm_logging_obj,
|
|
provider_config=provider_config,
|
|
api_base=api_base,
|
|
api_key=api_key,
|
|
client=client,
|
|
timeout=timeout,
|
|
headers=headers,
|
|
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
|
litellm_metadata=_build_litellm_metadata(kwargs),
|
|
query_params=query_params,
|
|
)
|
|
elif _custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS:
|
|
api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
|
# set API KEY
|
|
api_key = dynamic_api_key or litellm.api_key or litellm.openai_key or get_secret_str("AZURE_API_KEY")
|
|
|
|
api_version = api_version or litellm_params.api_version or "2024-10-01-preview"
|
|
|
|
realtime_protocol = (
|
|
kwargs.get("realtime_protocol")
|
|
or litellm_params.get("realtime_protocol")
|
|
or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL")
|
|
)
|
|
if realtime_protocol is None and (query_params or {}).get("intent") == "transcription":
|
|
realtime_protocol = "GA"
|
|
realtime_protocol = realtime_protocol or "beta"
|
|
resolved_azure_ad_token: Final = (
|
|
None if api_key else get_azure_ad_token(GenericLiteLLMParams(**kwargs, azure_ad_token=azure_ad_token))
|
|
)
|
|
await azure_realtime.async_realtime(
|
|
model=model,
|
|
websocket=websocket,
|
|
api_base=api_base,
|
|
api_key=api_key,
|
|
api_version=api_version,
|
|
azure_ad_token=resolved_azure_ad_token,
|
|
client=None,
|
|
timeout=timeout,
|
|
logging_obj=litellm_logging_obj,
|
|
realtime_protocol=realtime_protocol,
|
|
query_params=query_params,
|
|
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
|
litellm_metadata=_build_litellm_metadata(kwargs),
|
|
)
|
|
elif _custom_llm_provider == "openai":
|
|
api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or "https://api.openai.com/"
|
|
# set API KEY
|
|
api_key = dynamic_api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY")
|
|
|
|
await openai_realtime.async_realtime(
|
|
model=model,
|
|
websocket=websocket,
|
|
logging_obj=litellm_logging_obj,
|
|
api_base=api_base,
|
|
api_key=api_key,
|
|
client=None,
|
|
timeout=timeout,
|
|
query_params=query_params,
|
|
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
|
litellm_metadata=_build_litellm_metadata(kwargs),
|
|
)
|
|
elif _custom_llm_provider == "bedrock":
|
|
# Extract AWS parameters from kwargs
|
|
aws_region_name: Final = kwargs.get("aws_region_name")
|
|
aws_access_key_id: Final = kwargs.get("aws_access_key_id")
|
|
aws_secret_access_key: Final = kwargs.get("aws_secret_access_key")
|
|
aws_session_token: Final = kwargs.get("aws_session_token")
|
|
aws_role_name: Final = kwargs.get("aws_role_name")
|
|
aws_session_name: Final = kwargs.get("aws_session_name")
|
|
aws_profile_name: Final = kwargs.get("aws_profile_name")
|
|
aws_web_identity_token: Final = kwargs.get("aws_web_identity_token")
|
|
aws_sts_endpoint: Final = kwargs.get("aws_sts_endpoint")
|
|
aws_bedrock_runtime_endpoint: Final = kwargs.get("aws_bedrock_runtime_endpoint")
|
|
aws_external_id: Final = kwargs.get("aws_external_id")
|
|
|
|
await bedrock_realtime.async_realtime(
|
|
model=model,
|
|
websocket=websocket,
|
|
logging_obj=litellm_logging_obj,
|
|
api_base=dynamic_api_base or api_base,
|
|
api_key=dynamic_api_key or api_key,
|
|
timeout=timeout,
|
|
aws_region_name=aws_region_name,
|
|
aws_access_key_id=aws_access_key_id,
|
|
aws_secret_access_key=aws_secret_access_key,
|
|
aws_session_token=aws_session_token,
|
|
aws_role_name=aws_role_name,
|
|
aws_session_name=aws_session_name,
|
|
aws_profile_name=aws_profile_name,
|
|
aws_web_identity_token=aws_web_identity_token,
|
|
aws_sts_endpoint=aws_sts_endpoint,
|
|
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
|
|
aws_external_id=aws_external_id,
|
|
)
|
|
elif _custom_llm_provider == "xai":
|
|
api_base = (
|
|
dynamic_api_base or litellm_params.api_base or get_secret_str("XAI_API_BASE") or "https://api.x.ai/v1"
|
|
)
|
|
# set API KEY
|
|
api_key = XAIModelInfo.get_api_key(dynamic_api_key, legacy_generic_before_env=True)
|
|
|
|
await xai_realtime.async_realtime(
|
|
model=model,
|
|
websocket=websocket,
|
|
logging_obj=litellm_logging_obj,
|
|
api_base=api_base,
|
|
api_key=api_key,
|
|
client=None,
|
|
timeout=timeout,
|
|
query_params=query_params,
|
|
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
|
litellm_metadata=_build_litellm_metadata(kwargs),
|
|
)
|
|
elif _custom_llm_provider == "vertex_ai":
|
|
vertex_credentials: Final = (
|
|
kwargs.get("vertex_credentials")
|
|
or kwargs.get("vertex_ai_credentials")
|
|
or get_secret_str("VERTEXAI_CREDENTIALS")
|
|
)
|
|
vertex_project: Final = (
|
|
kwargs.get("vertex_project")
|
|
or kwargs.get("vertex_ai_project")
|
|
or litellm.vertex_project
|
|
or get_secret_str("VERTEXAI_PROJECT")
|
|
)
|
|
vertex_location: Final = (
|
|
kwargs.get("vertex_location")
|
|
or kwargs.get("vertex_ai_location")
|
|
or litellm.vertex_location
|
|
or get_secret_str("VERTEXAI_LOCATION")
|
|
)
|
|
|
|
resolved_location: Final = vertex_llm_base.get_vertex_region(vertex_region=vertex_location, model=model)
|
|
|
|
(
|
|
access_token,
|
|
resolved_project,
|
|
) = await _resolve_vertex_access_token_bounded(
|
|
credentials=vertex_credentials,
|
|
project_id=vertex_project,
|
|
resolver=vertex_access_token_resolver,
|
|
timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
|
)
|
|
|
|
vertex_realtime_config: Final = VertexAIRealtimeConfig(
|
|
access_token=access_token,
|
|
project=resolved_project,
|
|
location=resolved_location,
|
|
)
|
|
|
|
await base_llm_http_handler.async_realtime(
|
|
model=model,
|
|
websocket=websocket,
|
|
logging_obj=litellm_logging_obj,
|
|
provider_config=vertex_realtime_config,
|
|
api_base=dynamic_api_base or litellm_params.api_base,
|
|
api_key=None,
|
|
client=client,
|
|
timeout=timeout,
|
|
headers=headers,
|
|
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
|
litellm_metadata=_build_litellm_metadata(kwargs),
|
|
query_params=query_params,
|
|
)
|
|
else:
|
|
raise ValueError(f"Unsupported model: {model}")
|
|
|
|
|
|
def _is_transcription_only_realtime_model(model: str, custom_llm_provider: str) -> bool:
|
|
try:
|
|
model_info: Final = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
|
except Exception: # noqa: BLE001 # get_model_info raises bare Exception for unmapped models
|
|
return False
|
|
if model_info.get("mode") == "audio_transcription":
|
|
return True
|
|
return "/v1/realtime/transcription_sessions" in (model_info.get("supported_endpoints") or ())
|
|
|
|
|
|
_TRANSCRIPTION_QUERY_PARAMS: Final[RealtimeQueryParams] = {"intent": "transcription"}
|
|
|
|
|
|
def _azure_realtime_health_protocol(
|
|
model: str, realtime_protocol: str | None, model_params: Mapping[str, object]
|
|
) -> tuple[str, RealtimeQueryParams | None]:
|
|
query_params: Final = _TRANSCRIPTION_QUERY_PARAMS if _is_transcription_only_realtime_model(model, "azure") else None
|
|
configured_raw: Final = (
|
|
realtime_protocol or model_params.get("realtime_protocol") or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL")
|
|
)
|
|
configured: Final = configured_raw if isinstance(configured_raw, str) else None
|
|
if configured is not None:
|
|
return configured, query_params
|
|
if query_params is not None:
|
|
return "GA", query_params
|
|
return "beta", None
|
|
|
|
|
|
def _realtime_health_check_auth_headers(
|
|
custom_llm_provider: str, api_key: str | None, model_params: Mapping[str, Any]
|
|
) -> Mapping[str, str | None]:
|
|
if custom_llm_provider != "azure":
|
|
return MappingProxyType({"api-key": api_key})
|
|
return azure_realtime.get_auth_headers(
|
|
api_key=api_key,
|
|
azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))),
|
|
)
|
|
|
|
|
|
async def _realtime_health_check(
|
|
model: str,
|
|
custom_llm_provider: str,
|
|
api_key: str | None,
|
|
api_base: str | None = None,
|
|
api_version: str | None = None,
|
|
realtime_protocol: str | None = None,
|
|
model_params: dict | None = None,
|
|
):
|
|
"""
|
|
Health check for realtime API - tries connection to the realtime API websocket
|
|
|
|
Args:
|
|
model: str - model name
|
|
api_base: str - api base
|
|
api_version: Optional[str] - api version
|
|
api_key: str - api key
|
|
custom_llm_provider: str - custom llm provider
|
|
realtime_protocol: Optional[str] - protocol version ("GA"/"v1" for GA path, "beta" for beta path);
|
|
None resolves it for Azure from model_params/env, with transcription-only models probing GA
|
|
plus intent=transcription the way real calls do
|
|
|
|
Returns:
|
|
bool - True if connection is successful, False otherwise
|
|
Raises:
|
|
Exception - if the connection is not successful
|
|
"""
|
|
import websockets
|
|
|
|
url: str | None = None
|
|
auth_headers: Final = _realtime_health_check_auth_headers(
|
|
custom_llm_provider=custom_llm_provider,
|
|
api_key=api_key,
|
|
model_params=model_params or _EMPTY_MODEL_PARAMS,
|
|
)
|
|
if custom_llm_provider == "azure":
|
|
resolved_protocol, azure_query_params = _azure_realtime_health_protocol(
|
|
model=model,
|
|
realtime_protocol=realtime_protocol,
|
|
model_params=model_params or _EMPTY_MODEL_PARAMS,
|
|
)
|
|
url = azure_realtime._construct_url(
|
|
api_base=api_base or "",
|
|
model=model,
|
|
api_version=api_version or "2024-10-01-preview",
|
|
realtime_protocol=resolved_protocol,
|
|
query_params=azure_query_params,
|
|
)
|
|
elif custom_llm_provider == "openai":
|
|
url = openai_realtime._construct_url(
|
|
api_base=api_base or "https://api.openai.com/",
|
|
query_params={"model": model},
|
|
)
|
|
elif custom_llm_provider == "xai":
|
|
url = xai_realtime._construct_url(api_base=api_base or "https://api.x.ai/v1", query_params={"model": model})
|
|
elif custom_llm_provider == "vertex_ai":
|
|
vertex_model_params: Final = model_params or {}
|
|
resolved_location: Final = vertex_llm_base.get_vertex_region(
|
|
vertex_region=VertexBase.safe_get_vertex_ai_location(vertex_model_params),
|
|
model=model,
|
|
)
|
|
(
|
|
access_token,
|
|
resolved_project,
|
|
) = await _resolve_vertex_access_token_bounded(
|
|
credentials=VertexBase.safe_get_vertex_ai_credentials(vertex_model_params),
|
|
project_id=VertexBase.safe_get_vertex_ai_project(vertex_model_params),
|
|
resolver=vertex_access_token_resolver,
|
|
timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
|
)
|
|
vertex_realtime_config: Final = VertexAIRealtimeConfig(
|
|
access_token=access_token,
|
|
project=resolved_project,
|
|
location=resolved_location,
|
|
)
|
|
url = vertex_realtime_config.get_complete_url(api_base=api_base, model=model)
|
|
ssl_context = get_shared_realtime_ssl_context()
|
|
headers: Final = vertex_realtime_config.validate_environment(headers={}, model=model, api_key=None)
|
|
async with websockets.connect(
|
|
url,
|
|
additional_headers=headers,
|
|
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
|
ssl=ssl_context,
|
|
):
|
|
return True
|
|
else:
|
|
raise ValueError(f"Unsupported model: {model}")
|
|
ssl_context = get_shared_realtime_ssl_context()
|
|
async with websockets.connect(
|
|
url,
|
|
additional_headers=auth_headers,
|
|
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
|
ssl=ssl_context,
|
|
):
|
|
return True
|