Merge remote-tracking branch 'origin/main' into litellm_realtime_health_check_first_event

# Conflicts:
#	litellm/realtime_api/main.py
#	tests/test_litellm/realtime_api/test_main.py
This commit is contained in:
mateo-berri 2026-09-15 05:34:58 -07:00
commit 8880e0a03d
2 changed files with 102 additions and 22 deletions

View file

@ -34,6 +34,7 @@ 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
@ -62,6 +63,16 @@ _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
@ -673,34 +684,46 @@ async def _realtime_health_check(
"""
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=api_key,
model_params=model_params or _EMPTY_MODEL_PARAMS,
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=model_params or _EMPTY_MODEL_PARAMS,
model_params=resolved_params,
)
url = azure_realtime._construct_url(
api_base=api_base or "",
api_base=resolved_api_base or "",
model=model,
api_version=api_version or "2024-10-01-preview",
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=api_base or "https://api.openai.com/",
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=api_base or "https://api.x.ai/v1", query_params={"model": model})
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 = model_params or {}
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,
@ -719,19 +742,19 @@ async def _realtime_health_check(
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()
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=ssl_context,
ssl=vertex_ssl_context,
):
return True
else:
raise ValueError(f"Unsupported model: {model}")
ssl_context = get_shared_realtime_ssl_context()
ssl_context: Final = get_shared_realtime_ssl_context()
async with websockets.connect(
url,
additional_headers=auth_headers,

View file

@ -9,6 +9,7 @@ from websockets.exceptions import ConnectionClosedError
from websockets.frames import Close
import litellm
from litellm.models.credentials import CredentialItem
from litellm.realtime_api import main as realtime_main
from litellm.realtime_api.main import _with_resolved_session_model
@ -246,6 +247,72 @@ class _CapturingConnect:
return None
@pytest.mark.asyncio
async def test_azure_health_check_resolves_stored_credentials(monkeypatch):
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="azure-rt",
credential_values={
"api_key": "sk-from-credential",
"api_base": "https://example.openai.azure.com",
"api_version": "2025-04-01-preview",
},
credential_info={},
)
],
)
connect = _CapturingConnect()
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
model="gpt-realtime",
custom_llm_provider="azure",
api_key=None,
realtime_protocol="beta",
model_params={"model": "azure/gpt-realtime", "litellm_credential_name": "azure-rt"},
)
assert connect.kwargs["additional_headers"] == {"api-key": "sk-from-credential"}
assert connect.url is not None
assert connect.url.startswith("wss://example.openai.azure.com")
assert "api-version=2025-04-01-preview" in connect.url
@pytest.mark.asyncio
@pytest.mark.parametrize(
("custom_llm_provider", "model", "expected_url"),
[
("xai", "grok-voice-latest", "wss://api.x.ai/v1/realtime?model=grok-voice-latest"),
("openai", "gpt-realtime", "wss://api.openai.com/v1/realtime?model=gpt-realtime"),
],
)
async def test_bearer_health_check_sends_stored_credential_as_bearer_token(
monkeypatch, custom_llm_provider: str, model: str, expected_url: str
):
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="voice-key",
credential_values={"api_key": "sk-from-credential"},
credential_info={},
)
],
)
connect = _CapturingConnect(_ScriptedConnection(_OPENAI_SESSION_CREATED_EVENT))
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
model=model,
custom_llm_provider=custom_llm_provider,
api_key=None,
model_params={"model": f"{custom_llm_provider}/{model}", "litellm_credential_name": "voice-key"},
)
assert connect.kwargs["additional_headers"] == {"Authorization": "Bearer sk-from-credential"}
assert connect.url == expected_url
@pytest.mark.asyncio
async def test_azure_health_check_probes_ga_transcription_url_for_transcription_model(local_model_cost_map):
"""Regression for LIT-6240: transcription-only models (mode audio_transcription
@ -407,16 +474,6 @@ async def test_openai_health_check_reports_the_first_error_event_as_unhealthy(
assert connect.url == "wss://api.openai.com/v1/realtime?model=gpt-realtime"
@pytest.mark.asyncio
async def test_openai_health_check_sends_the_api_key_as_a_bearer_token():
connect: Final = _CapturingConnect(_ScriptedConnection(_OPENAI_SESSION_CREATED_EVENT))
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
model="gpt-realtime", custom_llm_provider="openai", api_key="sk-real"
)
assert connect.kwargs["additional_headers"] == {"Authorization": "Bearer sk-real"}
@pytest.mark.asyncio
async def test_openai_health_check_without_an_api_key_sends_no_auth_header_and_is_unhealthy():
connect: Final = _CapturingConnect(_ScriptedConnection(_OPENAI_MISSING_AUTH_EVENT))