mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge 5ff9468cc9 into 4ece6c9fb8
This commit is contained in:
commit
a48162f113
3 changed files with 99 additions and 17 deletions
|
|
@ -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/<name>``, 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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue