diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 923a6cc5743..5573e240449 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -3682,7 +3682,7 @@ def _build_vertex_live_setup_model_rewriter( return rewrite -def _resolve_alias_to_upstream_model(setup_model: str, llm_router: Router | None) -> str: +def resolve_alias_to_upstream_model(setup_model: str, llm_router: Router | None) -> str: """ The Live SDK wraps whatever the caller typed as ``models/``, so a gateway alias arrives prefixed """ @@ -3706,6 +3706,9 @@ def _resolve_alias_to_upstream_model(setup_model: str, llm_router: Router | None return upstream.removeprefix(f"{provider}/") +_resolve_alias_to_upstream_model: Final = resolve_alias_to_upstream_model + + async def vertex_ai_live_websocket_passthrough( websocket: WebSocket, model: str | None = None, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index a5b414e5a18..44eecc7f4bf 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -2910,7 +2910,32 @@ async def _relay_passthrough_response_bytes( bind_budget_reservation_to_callbacks(logging_obj.litellm_params) -def _extract_model_from_vertex_ai_setup(setup_response: Mapping[str, object]) -> str | None: +def _get_proxy_router() -> litellm.Router | None: + try: + from litellm.proxy.proxy_server import llm_router + except (ImportError, AttributeError): + return None + else: + return llm_router + + +def _extract_model_path_from_setup(setup_response: Mapping[str, object]) -> str | None: + if isinstance(setup_response, dict): + direct_model: Final = setup_response.get("model") + if isinstance(direct_model, str): + return direct_model + setup_obj: Final = setup_response.get("setup") + if isinstance(setup_obj, dict): + nested_model: Final = setup_obj.get("model") + if isinstance(nested_model, str): + return nested_model + return None + + +def _extract_model_from_vertex_ai_setup( + setup_response: Mapping[str, object], + llm_router: litellm.Router | None = None, +) -> str | None: """ Extract the model name from Vertex AI Live setup response. @@ -2921,22 +2946,26 @@ def _extract_model_from_vertex_ai_setup(setup_response: Mapping[str, object]) -> We extract just the model name: "gemini-2.0-flash-live-preview-04-09" """ try: - # Handle both direct model field and nested setup.model field - model_path = None - if isinstance(setup_response, dict): - if "model" in setup_response: - model_path = setup_response["model"] - elif ( - "setup" in setup_response - and isinstance(setup_response["setup"], dict) - and "model" in setup_response["setup"] - ): - model_path = setup_response["setup"]["model"] + model_path: Final = _extract_model_path_from_setup(setup_response) - if isinstance(model_path, str) and "/models/" in model_path: - # Extract the model name after the last "/models/" - model_name: Final = model_path.split("/models/")[-1] - return model_name + if isinstance(model_path, str): + if "/models/" in model_path: + # Extract the model name after the last "/models/" + model_name: Final = model_path.split("/models/")[-1] + return model_name + + active_router: Final[litellm.Router | None] = llm_router if llm_router is not None else _get_proxy_router() + + if active_router is not None: + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + resolve_alias_to_upstream_model, + ) + + resolved: Final = resolve_alias_to_upstream_model(model_path, active_router) + if resolved != model_path: + return resolved + + return model_path.rsplit("/", 1)[-1] except Exception as e: verbose_proxy_logger.debug("Error extracting model from setup response: %s", e) return None diff --git a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py index 8b3dc436b8f..5e2bf8a541b 100644 --- a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py +++ b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py @@ -996,5 +996,55 @@ class TestVertexAILivePassthroughErrorHandling: assert "kwargs" in result +def test_extract_model_from_vertex_ai_setup_with_alias() -> None: + from typing import Final + + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _extract_model_from_vertex_ai_setup, + ) + + mock_router: Final = MagicMock() + mock_router.get_model_list.return_value = [ + { + "model_name": "transcribe_live", + "litellm_params": {"model": "vertex_ai/gemini-3.5-transcribe-live-preview"}, + } + ] + + full_path_setup: Final = { + "setup": { + "model": "projects/my-proj/locations/us-central1/publishers/google/models/gemini-2.0-flash-exp" + } + } + assert _extract_model_from_vertex_ai_setup(full_path_setup) == "gemini-2.0-flash-exp" + + alias_setup: Final = {"setup": {"model": "transcribe_live"}} + assert ( + _extract_model_from_vertex_ai_setup(alias_setup, llm_router=mock_router) + == "gemini-3.5-transcribe-live-preview" + ) + + prefixed_alias_setup: Final = {"setup": {"model": "models/transcribe_live"}} + assert ( + _extract_model_from_vertex_ai_setup(prefixed_alias_setup, llm_router=mock_router) + == "gemini-3.5-transcribe-live-preview" + ) + + direct_setup: Final = { + "model": "projects/my-proj/locations/us-central1/publishers/google/models/gemini-2.0-flash-exp" + } + assert _extract_model_from_vertex_ai_setup(direct_setup) == "gemini-2.0-flash-exp" + + unaliased_setup: Final = {"setup": {"model": "gemini-2.0-flash-exp"}} + assert _extract_model_from_vertex_ai_setup(unaliased_setup) == "gemini-2.0-flash-exp" + assert ( + _extract_model_from_vertex_ai_setup(unaliased_setup, llm_router=mock_router) + == "gemini-2.0-flash-exp" + ) + + empty_setup: Final = {} + assert _extract_model_from_vertex_ai_setup(empty_setup) is None + + if __name__ == "__main__": pytest.main([__file__])