mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-23 00:41:40 +00:00
Merge origin/main (a9ee15372f) into the typed AwsAuthParams refactor so the
session tags PR #40446 added land in the struct: resolve_credentials
canonicalizes aws_session_tags before STS, the realtime path forwards them,
and Files upload/download plus bodiless S3 signing now assume the role with
the tags instead of dropping them.
725 lines
28 KiB
Python
725 lines
28 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.credential_accessor import CredentialAccessor
|
|
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, azure_realtime_protocol_for_client
|
|
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({})
|
|
_EMPTY_AUTH_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
|
|
|
|
|
|
def _model_params_with_stored_credentials(model_params: Mapping[str, Any]) -> Mapping[str, Any]:
|
|
credential_name: Final = model_params.get("litellm_credential_name")
|
|
credential_values: Final = (
|
|
CredentialAccessor.get_credential_values(credential_name)
|
|
if isinstance(credential_name, str)
|
|
else _EMPTY_MODEL_PARAMS
|
|
)
|
|
return MappingProxyType({**credential_values, **model_params})
|
|
|
|
|
|
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"
|
|
|
|
configured_realtime_protocol: Final = (
|
|
kwargs.get("realtime_protocol")
|
|
or litellm_params.get("realtime_protocol")
|
|
or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL")
|
|
)
|
|
realtime_protocol: Final = azure_realtime_protocol_for_client(
|
|
configured_realtime_protocol, query_params=query_params, websocket=websocket
|
|
)
|
|
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")
|
|
aws_session_tags: Final = kwargs.get("aws_session_tags")
|
|
|
|
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,
|
|
aws_session_tags=aws_session_tags,
|
|
)
|
|
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
|
|
return "GA", query_params
|
|
|
|
|
|
def _realtime_health_check_auth_headers(
|
|
custom_llm_provider: str, api_key: str | None, model_params: Mapping[str, Any]
|
|
) -> Mapping[str, str]:
|
|
if custom_llm_provider == "azure":
|
|
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))),
|
|
)
|
|
if api_key is None:
|
|
return _EMPTY_AUTH_HEADERS
|
|
return MappingProxyType({"Authorization": f"Bearer {api_key}"})
|
|
|
|
|
|
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 and otherwise probes GA, the upstream a client
|
|
without the OpenAI-Beta header is bridged to, with transcription-only models adding intent=transcription
|
|
|
|
Returns:
|
|
bool - True if connection is successful, False otherwise
|
|
Raises:
|
|
Exception - if the connection is not successful
|
|
"""
|
|
import websockets
|
|
|
|
resolved_params: Final = _model_params_with_stored_credentials(model_params or _EMPTY_MODEL_PARAMS)
|
|
resolved_api_key: Final = cast( # cast-ok: provider parameters expose optional string credentials
|
|
str | None, api_key or resolved_params.get("api_key")
|
|
)
|
|
resolved_api_base: Final = cast( # cast-ok: provider parameters expose optional string endpoints
|
|
str | None, api_base or resolved_params.get("api_base")
|
|
)
|
|
resolved_api_version: Final = cast( # cast-ok: provider parameters expose optional string versions
|
|
str | None, api_version or resolved_params.get("api_version")
|
|
)
|
|
url: str | None = None
|
|
auth_headers: Final = _realtime_health_check_auth_headers(
|
|
custom_llm_provider=custom_llm_provider,
|
|
api_key=resolved_api_key,
|
|
model_params=resolved_params,
|
|
)
|
|
if custom_llm_provider == "azure":
|
|
resolved_protocol, azure_query_params = _azure_realtime_health_protocol(
|
|
model=model,
|
|
realtime_protocol=realtime_protocol,
|
|
model_params=resolved_params,
|
|
)
|
|
url = azure_realtime._construct_url(
|
|
api_base=resolved_api_base or "",
|
|
model=model,
|
|
api_version=resolved_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=resolved_api_base or "https://api.openai.com/",
|
|
query_params={"model": model},
|
|
)
|
|
elif custom_llm_provider == "xai":
|
|
url = xai_realtime._construct_url(
|
|
api_base=resolved_api_base or "https://api.x.ai/v1", query_params={"model": model}
|
|
)
|
|
elif custom_llm_provider == "vertex_ai":
|
|
vertex_model_params: Final = dict(resolved_params)
|
|
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=resolved_api_base, model=model)
|
|
vertex_ssl_context: Final = 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=vertex_ssl_context,
|
|
):
|
|
return True
|
|
else:
|
|
raise ValueError(f"Unsupported model: {model}")
|
|
ssl_context: Final = 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
|