mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
Merge 4b7a97fed1 into 02dcc4d347
This commit is contained in:
commit
4c24940b7e
2 changed files with 99 additions and 4 deletions
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
import asyncio
|
||||
import os
|
||||
import ssl
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Literal, cast
|
||||
|
|
@ -596,6 +597,20 @@ def _realtime_health_check_auth_headers(
|
|||
)
|
||||
|
||||
|
||||
def _get_realtime_ssl_context(url: str | None) -> bool | str | ssl.SSLContext | None:
|
||||
"""SSL value for websockets.connect, matching OpenAIRealtime._get_ssl_config.
|
||||
|
||||
websockets rejects both an ssl= argument on a ws:// URI and ssl=False, so
|
||||
ws:// gets None and a disabled-verification context becomes True.
|
||||
"""
|
||||
if url and url.startswith("ws://"):
|
||||
return None
|
||||
ssl_config = get_shared_realtime_ssl_context()
|
||||
if ssl_config is False:
|
||||
return True
|
||||
return ssl_config
|
||||
|
||||
|
||||
async def _realtime_health_check(
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
|
|
@ -672,7 +687,7 @@ async def _realtime_health_check(
|
|||
location=resolved_location,
|
||||
)
|
||||
url = vertex_realtime_config.get_complete_url(api_base=api_base, model=model)
|
||||
ssl_context = get_shared_realtime_ssl_context()
|
||||
ssl_context = _get_realtime_ssl_context(url)
|
||||
headers: Final = vertex_realtime_config.validate_environment(headers={}, model=model, api_key=None)
|
||||
async with websockets.connect(
|
||||
url,
|
||||
|
|
@ -683,7 +698,7 @@ async def _realtime_health_check(
|
|||
return True
|
||||
else:
|
||||
raise ValueError(f"Unsupported model: {model}")
|
||||
ssl_context = get_shared_realtime_ssl_context()
|
||||
ssl_context = _get_realtime_ssl_context(url)
|
||||
async with websockets.connect(
|
||||
url,
|
||||
additional_headers=auth_headers,
|
||||
|
|
|
|||
|
|
@ -1,14 +1,15 @@
|
|||
import asyncio
|
||||
import ssl as ssl_module
|
||||
import time
|
||||
from types import TracebackType
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.realtime_api import main as realtime_main
|
||||
from litellm.realtime_api.main import _with_resolved_session_model
|
||||
from litellm.realtime_api.main import _get_realtime_ssl_context, _with_resolved_session_model
|
||||
|
||||
|
||||
class FakeLogging:
|
||||
|
|
@ -197,9 +198,11 @@ def test_client_secret_forwards_nested_transcription_model_untouched(monkeypatch
|
|||
class _CapturingConnect:
|
||||
def __init__(self) -> None:
|
||||
self.url: str | None = None
|
||||
self.kwargs: dict[str, object] = {}
|
||||
|
||||
def __call__(self, url: str, **kwargs: object) -> "_CapturingConnect":
|
||||
self.url = url
|
||||
self.kwargs = kwargs
|
||||
return self
|
||||
|
||||
async def __aenter__(self) -> MagicMock:
|
||||
|
|
@ -294,3 +297,80 @@ async def test_azure_health_check_honors_deployment_realtime_protocol():
|
|||
model_params={"realtime_protocol": "GA"},
|
||||
)
|
||||
assert connect.url == "wss://my-endpoint.openai.azure.com/openai/v1/realtime?model=gpt-4o-realtime-preview"
|
||||
|
||||
|
||||
def test_realtime_ssl_context_is_none_for_ws():
|
||||
assert _get_realtime_ssl_context("ws://host/v1/realtime") is None
|
||||
|
||||
|
||||
def test_realtime_ssl_context_normalizes_false_to_true_for_wss(monkeypatch):
|
||||
"""ssl_verify=False makes the shared context False, which websockets rejects."""
|
||||
monkeypatch.setattr(realtime_main, "get_shared_realtime_ssl_context", lambda: False)
|
||||
assert _get_realtime_ssl_context("wss://host/v1/realtime") is True
|
||||
|
||||
|
||||
def test_realtime_ssl_context_passes_through_context_for_wss(monkeypatch):
|
||||
ctx = ssl_module.create_default_context()
|
||||
monkeypatch.setattr(realtime_main, "get_shared_realtime_ssl_context", lambda: ctx)
|
||||
assert _get_realtime_ssl_context("wss://host/v1/realtime") is ctx
|
||||
|
||||
|
||||
def test_realtime_ssl_context_defaults_to_wss_branch_when_url_is_none(monkeypatch):
|
||||
monkeypatch.setattr(realtime_main, "get_shared_realtime_ssl_context", lambda: True)
|
||||
assert _get_realtime_ssl_context(None) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_omits_ssl_for_ws_url():
|
||||
"""Regression for #31613: websockets rejects ssl= on a ws:// URI, which broke
|
||||
health checks against OpenAI-compatible servers behind an http:// api_base."""
|
||||
connect = _CapturingConnect()
|
||||
with patch("websockets.connect", connect):
|
||||
assert await realtime_main._realtime_health_check(
|
||||
model="some-model",
|
||||
custom_llm_provider="openai",
|
||||
api_key="fake-key",
|
||||
api_base="http://localhost:8030/v1",
|
||||
)
|
||||
assert connect.url is not None and connect.url.startswith("ws://")
|
||||
assert connect.kwargs["ssl"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_keeps_ssl_for_wss_url():
|
||||
connect = _CapturingConnect()
|
||||
with patch("websockets.connect", connect):
|
||||
assert await realtime_main._realtime_health_check(
|
||||
model="some-model",
|
||||
custom_llm_provider="openai",
|
||||
api_key="fake-key",
|
||||
api_base="https://api.openai.com/",
|
||||
)
|
||||
assert connect.url is not None and connect.url.startswith("wss://")
|
||||
assert connect.kwargs["ssl"] is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_vertex_branch_routes_through_ssl_helper(monkeypatch):
|
||||
connect = _CapturingConnect()
|
||||
fake_config = MagicMock()
|
||||
fake_config.get_complete_url.return_value = "wss://vertex.example/v1/realtime"
|
||||
fake_config.validate_environment.return_value = {}
|
||||
monkeypatch.setattr(
|
||||
realtime_main.vertex_llm_base, "get_vertex_region", lambda *args, **kwargs: "us-central1"
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
realtime_main.vertex_llm_base,
|
||||
"_ensure_access_token_async",
|
||||
AsyncMock(return_value=("token", "project")),
|
||||
)
|
||||
monkeypatch.setattr(realtime_main, "VertexAIRealtimeConfig", lambda **kwargs: fake_config)
|
||||
with patch("websockets.connect", connect):
|
||||
assert await realtime_main._realtime_health_check(
|
||||
model="gemini-realtime",
|
||||
custom_llm_provider="vertex_ai",
|
||||
api_key=None,
|
||||
api_base="https://vertex.example",
|
||||
)
|
||||
assert connect.url == "wss://vertex.example/v1/realtime"
|
||||
assert connect.kwargs["ssl"] is not None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue