mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
Merge pull request #32077 from BerriAI/litellm_realtime_internal_user_access
fix(proxy): route realtime HTTP endpoints through router for credenti…
This commit is contained in:
commit
8f6abb5565
5 changed files with 399 additions and 21 deletions
|
|
@ -435,9 +435,6 @@ async def route_request(
|
|||
"aget_run",
|
||||
"acancel_run",
|
||||
"adelete_run",
|
||||
"acreate_realtime_client_secret",
|
||||
"arealtime_calls",
|
||||
"acreate_realtime_transcription_session",
|
||||
]:
|
||||
# If a model is provided, get its credentials from the router
|
||||
model = data.get("model")
|
||||
|
|
|
|||
|
|
@ -40,6 +40,12 @@ vertex_llm_base = VertexBase()
|
|||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
|
||||
|
||||
def _with_resolved_session_model(session: dict[str, Any], model_name: str) -> dict[str, Any]:
|
||||
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: dict = {**(kwargs.get("litellm_metadata") or {})}
|
||||
|
|
@ -134,6 +140,8 @@ async def acreate_realtime_client_secret(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
request_data = 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,
|
||||
|
|
@ -249,6 +257,8 @@ async def arealtime_calls(
|
|||
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,
|
||||
|
|
|
|||
|
|
@ -1118,6 +1118,13 @@ class Router:
|
|||
self.adelete_responses = self.factory_function(litellm.adelete_responses, call_type="adelete_responses")
|
||||
self.alist_input_items = self.factory_function(litellm.alist_input_items, call_type="alist_input_items")
|
||||
self._arealtime = self.factory_function(litellm._arealtime, call_type="_arealtime")
|
||||
self.acreate_realtime_client_secret = self.factory_function(
|
||||
litellm.acreate_realtime_client_secret, call_type="acreate_realtime_client_secret"
|
||||
)
|
||||
self.arealtime_calls = self.factory_function(litellm.arealtime_calls, call_type="arealtime_calls")
|
||||
self.acreate_realtime_transcription_session = self.factory_function(
|
||||
litellm.acreate_realtime_transcription_session, call_type="acreate_realtime_transcription_session"
|
||||
)
|
||||
self._aresponses_websocket = self.factory_function(
|
||||
litellm._aresponses_websocket, call_type="_aresponses_websocket"
|
||||
)
|
||||
|
|
@ -5350,6 +5357,9 @@ class Router:
|
|||
"afile_delete",
|
||||
"afile_content",
|
||||
"_arealtime",
|
||||
"acreate_realtime_client_secret",
|
||||
"arealtime_calls",
|
||||
"acreate_realtime_transcription_session",
|
||||
"_aresponses_websocket",
|
||||
"acreate_fine_tuning_job",
|
||||
"acancel_fine_tuning_job",
|
||||
|
|
@ -5589,6 +5599,16 @@ class Router:
|
|||
original_function=original_function,
|
||||
**kwargs,
|
||||
)
|
||||
elif call_type in (
|
||||
"acreate_realtime_client_secret",
|
||||
"arealtime_calls",
|
||||
"acreate_realtime_transcription_session",
|
||||
):
|
||||
return await self._ageneric_api_call_with_fallbacks(
|
||||
original_function=original_function,
|
||||
client=client,
|
||||
**kwargs,
|
||||
)
|
||||
elif call_type in (
|
||||
"anthropic_messages",
|
||||
"_arealtime",
|
||||
|
|
|
|||
|
|
@ -3,9 +3,7 @@ import sys
|
|||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path
|
||||
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
|
@ -169,9 +167,7 @@ async def test_route_request_no_model_required_with_router_settings_and_no_route
|
|||
"messages": [{"role": "user", "content": "what llm are you"}],
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
litellm, "acompletion", return_value="fake_response"
|
||||
) as mock_completion:
|
||||
with patch.object(litellm, "acompletion", return_value="fake_response") as mock_completion:
|
||||
await route_request(data, None, "gpt-3.5-turbo", "acompletion")
|
||||
|
||||
mock_completion.assert_called_once_with(**data)
|
||||
|
|
@ -209,9 +205,7 @@ async def test_route_request_with_router_settings_override():
|
|||
assert call_kwargs["fallbacks"] == [{"gpt-3.5-turbo": ["gpt-4"]}]
|
||||
assert call_kwargs["num_retries"] == 5
|
||||
assert call_kwargs["timeout"] == 30
|
||||
assert call_kwargs["model_group_retry_policy"] == {
|
||||
"gpt-3.5-turbo": {"RateLimitErrorRetries": 3}
|
||||
}
|
||||
assert call_kwargs["model_group_retry_policy"] == {"gpt-3.5-turbo": {"RateLimitErrorRetries": 3}}
|
||||
# Verify unsupported settings were NOT merged
|
||||
assert "routing_strategy" not in call_kwargs
|
||||
assert "model_group_alias" not in call_kwargs
|
||||
|
|
@ -292,9 +286,7 @@ def test_mock_testing_kwarg_names_matches_dataclass():
|
|||
from litellm.proxy.route_llm_request import _MOCK_TESTING_KWARG_NAMES
|
||||
from litellm.types.router import MockRouterTestingParams
|
||||
|
||||
assert set(_MOCK_TESTING_KWARG_NAMES) == {
|
||||
f.name for f in fields(MockRouterTestingParams)
|
||||
}
|
||||
assert set(_MOCK_TESTING_KWARG_NAMES) == {f.name for f in fields(MockRouterTestingParams)}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -329,9 +321,7 @@ async def test_route_request_strips_mock_testing_flags(mock_flag):
|
|||
assert mock_flag not in data
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type", ["agenerate_content", "agenerate_content_stream"]
|
||||
)
|
||||
@pytest.mark.parametrize("route_type", ["agenerate_content", "agenerate_content_stream"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_maps_generation_config_for_google_routes(route_type):
|
||||
"""For Google generate_content routes, route_request must rename
|
||||
|
|
@ -359,9 +349,7 @@ async def test_route_request_maps_generation_config_for_google_routes(route_type
|
|||
assert call_kwargs["config"]["imageConfig"]["imageSize"] == "4K"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type", ["agenerate_content", "agenerate_content_stream"]
|
||||
)
|
||||
@pytest.mark.parametrize("route_type", ["agenerate_content", "agenerate_content_stream"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_preserves_existing_config_for_google_routes(route_type):
|
||||
"""If the caller already supplies `config`, route_request must not
|
||||
|
|
@ -379,3 +367,261 @@ async def test_route_request_preserves_existing_config_for_google_routes(route_t
|
|||
|
||||
call_kwargs = getattr(llm_router, route_type).call_args[1]
|
||||
assert call_kwargs["config"] == {"existing": True}
|
||||
|
||||
|
||||
async def _invoke_realtime_route(
|
||||
data: dict,
|
||||
llm_router,
|
||||
route_type: str = "acreate_realtime_client_secret",
|
||||
):
|
||||
llm_call = await route_request(data, llm_router, None, route_type)
|
||||
return await llm_call
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def openai_realtime_credential():
|
||||
import litellm
|
||||
from litellm.types.utils import CredentialItem
|
||||
|
||||
litellm.credential_list = [
|
||||
CredentialItem(
|
||||
credential_name="openai-realtime-cred",
|
||||
credential_info={"custom_llm_provider": "openai"},
|
||||
credential_values={"api_key": "resolved-credential-key"},
|
||||
)
|
||||
]
|
||||
yield
|
||||
litellm.credential_list = []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_realtime_wildcard_model_resolves_credentials(
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
POST /realtime/client_secrets with a request model like openai/gpt-realtime
|
||||
must match an openai/* deployment and forward its api_key upstream.
|
||||
"""
|
||||
import httpx
|
||||
import litellm
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
"api_key": "wildcard-realtime-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"litellm.realtime_api.main.base_llm_http_handler.async_realtime_client_secret_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handler:
|
||||
mock_handler.return_value = httpx.Response(200, json={"value": "ephemeral"})
|
||||
await _invoke_realtime_route(
|
||||
{"model": "openai/gpt-realtime"},
|
||||
router,
|
||||
)
|
||||
|
||||
assert mock_handler.call_args.kwargs["api_key"] == "wildcard-realtime-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_realtime_team_scoped_model_resolves_credentials(
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
Team-scoped deployments (team_public_model_name) must be selected when
|
||||
user_api_key_team_id is present, same as /chat/completions.
|
||||
"""
|
||||
import httpx
|
||||
import litellm
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "internal-realtime",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-realtime",
|
||||
"api_key": "team-realtime-key",
|
||||
},
|
||||
"model_info": {
|
||||
"team_id": "team-a",
|
||||
"team_public_model_name": "team-realtime",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"litellm.realtime_api.main.base_llm_http_handler.async_realtime_client_secret_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handler:
|
||||
mock_handler.return_value = httpx.Response(200, json={"value": "ephemeral"})
|
||||
await _invoke_realtime_route(
|
||||
{
|
||||
"model": "team-realtime",
|
||||
"metadata": {"user_api_key_team_id": "team-a"},
|
||||
},
|
||||
router,
|
||||
)
|
||||
|
||||
assert mock_handler.call_args.kwargs["api_key"] == "team-realtime-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_realtime_litellm_credential_name_resolves_api_key(
|
||||
openai_realtime_credential,
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
litellm_credential_name on a wildcard deployment must resolve to the stored
|
||||
api_key when routing acreate_realtime_client_secret through the router.
|
||||
"""
|
||||
import httpx
|
||||
import litellm
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
"litellm_credential_name": "openai-realtime-cred",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"litellm.realtime_api.main.base_llm_http_handler.async_realtime_client_secret_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handler:
|
||||
mock_handler.return_value = httpx.Response(200, json={"value": "ephemeral"})
|
||||
await _invoke_realtime_route({"model": "openai/gpt-realtime"}, router)
|
||||
|
||||
assert mock_handler.call_args.kwargs["api_key"] == "resolved-credential-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_realtime_unresolvable_model_raises_not_found(
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
An unknown model must not silently fall through to litellm with an empty
|
||||
OPENAI_API_KEY env var.
|
||||
"""
|
||||
import litellm
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
|
||||
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "other-model",
|
||||
"litellm_params": {"model": "openai/gpt-4", "api_key": "other-key"},
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"litellm.realtime_api.main.base_llm_http_handler.async_realtime_client_secret_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handler:
|
||||
with pytest.raises(ProxyModelNotFoundError):
|
||||
await _invoke_realtime_route({"model": "nonexistent-realtime-model"}, router)
|
||||
|
||||
mock_handler.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_realtime_calls_resolves_api_base(monkeypatch):
|
||||
"""
|
||||
/realtime/calls must resolve the deployment's api_base through the router so a
|
||||
non-default (self-hosted / proxied) OpenAI endpoint is honored, instead of
|
||||
defaulting to https://api.openai.com.
|
||||
"""
|
||||
import httpx
|
||||
import litellm
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "my-realtime",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-realtime",
|
||||
"api_key": "calls-key",
|
||||
"api_base": "https://custom-realtime.example.com/v1",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"litellm.realtime_api.main.base_llm_http_handler.async_realtime_calls_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handler:
|
||||
mock_handler.return_value = httpx.Response(200, content=b"v=0\r\n")
|
||||
await _invoke_realtime_route(
|
||||
{
|
||||
"model": "my-realtime",
|
||||
"openai_ephemeral_key": "ek_test",
|
||||
"sdp_body": b"v=0\r\n",
|
||||
},
|
||||
router,
|
||||
route_type="arealtime_calls",
|
||||
)
|
||||
|
||||
assert mock_handler.call_args.kwargs["api_base"] == "https://custom-realtime.example.com/v1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_realtime_transcription_session_resolves_credentials(monkeypatch):
|
||||
"""
|
||||
/realtime/transcription_sessions must resolve credentials through the router
|
||||
(wildcard deployment) rather than falling back to an empty OPENAI_API_KEY.
|
||||
"""
|
||||
import httpx
|
||||
import litellm
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {
|
||||
"model": "openai/*",
|
||||
"api_key": "transcription-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"litellm.realtime_api.main.base_llm_http_handler.async_realtime_transcription_session_handler",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handler:
|
||||
mock_handler.return_value = httpx.Response(200, json={"client_secret": {"value": "ephemeral"}})
|
||||
await _invoke_realtime_route(
|
||||
{"model": "openai/gpt-realtime"},
|
||||
router,
|
||||
route_type="acreate_realtime_transcription_session",
|
||||
)
|
||||
|
||||
assert mock_handler.call_args.kwargs["api_key"] == "transcription-key"
|
||||
|
|
|
|||
105
tests/test_litellm/realtime_api/test_main.py
Normal file
105
tests/test_litellm/realtime_api/test_main.py
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.realtime_api import main as realtime_main
|
||||
from litellm.realtime_api.main import _with_resolved_session_model
|
||||
|
||||
|
||||
class FakeLogging:
|
||||
def update_from_kwargs(self, **kwargs):
|
||||
pass
|
||||
|
||||
|
||||
def test_resolves_top_level_session_model():
|
||||
resolved = _with_resolved_session_model({"model": "alias/gpt-realtime"}, "gpt-realtime")
|
||||
assert resolved == {"model": "gpt-realtime"}
|
||||
|
||||
|
||||
def test_session_without_model_is_returned_unchanged():
|
||||
session = {"type": "realtime", "audio": {"input": {}}}
|
||||
assert _with_resolved_session_model(session, "gpt-realtime") == session
|
||||
|
||||
|
||||
def test_does_not_clobber_flat_transcription_model():
|
||||
"""The nested transcription model is a different model than the realtime
|
||||
conversation model and must not be overwritten with the routing model."""
|
||||
resolved = _with_resolved_session_model(
|
||||
{"model": "gpt-4o-realtime-preview", "input_audio_transcription": {"model": "whisper-1"}},
|
||||
"gpt-4o-realtime-preview",
|
||||
)
|
||||
assert resolved["input_audio_transcription"]["model"] == "whisper-1"
|
||||
|
||||
|
||||
def test_does_not_clobber_nested_audio_transcription_model():
|
||||
resolved = _with_resolved_session_model(
|
||||
{
|
||||
"model": "gpt-4o-realtime-preview",
|
||||
"audio": {"input": {"transcription": {"model": "whisper-1"}}},
|
||||
},
|
||||
"gpt-4o-realtime-preview",
|
||||
)
|
||||
assert resolved["audio"]["input"]["transcription"]["model"] == "whisper-1"
|
||||
|
||||
|
||||
def test_original_session_is_not_mutated():
|
||||
session = {"model": "alias/gpt-realtime"}
|
||||
_with_resolved_session_model(session, "gpt-realtime")
|
||||
assert session == {"model": "alias/gpt-realtime"}
|
||||
|
||||
|
||||
def _run_client_secret(session, model, monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def mock_handler(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return object()
|
||||
|
||||
def mock_get_llm_provider(model, api_base, api_key):
|
||||
return model, "openai", None, api_base
|
||||
|
||||
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
|
||||
monkeypatch.setattr(
|
||||
realtime_main.base_llm_http_handler,
|
||||
"async_realtime_client_secret_handler",
|
||||
mock_handler,
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
realtime_main.acreate_realtime_client_secret.__wrapped__(
|
||||
model=model,
|
||||
session=session,
|
||||
litellm_logging_obj=FakeLogging(),
|
||||
)
|
||||
)
|
||||
return captured
|
||||
|
||||
|
||||
def test_client_secret_session_model_takes_priority_over_top_level(monkeypatch):
|
||||
"""Backwards-compatible ordering: an explicit session.model wins over the
|
||||
top-level model, matching the proxy's own resolution order."""
|
||||
captured = _run_client_secret(
|
||||
session={"model": "gpt-realtime-session"},
|
||||
model="gpt-realtime-top-level",
|
||||
monkeypatch=monkeypatch,
|
||||
)
|
||||
assert captured["model"] == "gpt-realtime-session"
|
||||
assert captured["request_data"]["session"]["model"] == "gpt-realtime-session"
|
||||
|
||||
|
||||
def test_client_secret_forwards_nested_transcription_model_untouched(monkeypatch):
|
||||
captured = _run_client_secret(
|
||||
session={
|
||||
"model": "gpt-4o-realtime-preview",
|
||||
"input_audio_transcription": {"model": "whisper-1"},
|
||||
},
|
||||
model=None,
|
||||
monkeypatch=monkeypatch,
|
||||
)
|
||||
session = captured["request_data"]["session"]
|
||||
assert session["model"] == "gpt-4o-realtime-preview"
|
||||
assert session["input_audio_transcription"]["model"] == "whisper-1"
|
||||
Loading…
Add table
Reference in a new issue