mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(vertex): forward realtime health check params (#32550)
* fix(vertex): forward realtime health check params * refactor(vertex): resolve realtime health check params via VertexBase helpers Address review feedback on the vertex param forwarding: pass model_params through to _realtime_health_check and resolve vertex credentials, project, and location inside the vertex_ai branch using the existing VertexBase.safe_get_vertex_ai_* helpers, so provider-specific key extraction no longer lives in litellm_core_utils and dict-typed vertex_credentials are supported * test(vertex): move realtime health check test to mapped unit test path codecov/patch reported the vertex branch of _realtime_health_check as uncovered because tests/litellm_utils_tests is not part of the unit test groups that upload coverage. Move the test into tests/test_litellm/litellm_core_utils/test_health_check_helpers.py, which the core-utils group runs, keeping the same end-to-end assertions through litellm.ahealth_check --------- Co-authored-by: Aleksandr Liadov <72351793+AleksandrLiadov@users.noreply.github.com>
This commit is contained in:
parent
4e6ec995e7
commit
5aa5b8ef32
3 changed files with 76 additions and 4 deletions
|
|
@ -199,6 +199,7 @@ class HealthCheckHelpers:
|
|||
api_base=model_params.get("api_base", None),
|
||||
api_key=model_params.get("api_key", None),
|
||||
api_version=model_params.get("api_version", None),
|
||||
model_params=model_params,
|
||||
),
|
||||
"batch": lambda: HealthCheckHelpers._batch_health_check(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -515,6 +515,7 @@ async def _realtime_health_check(
|
|||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
realtime_protocol: Optional[str] = None,
|
||||
model_params: Optional[dict] = None,
|
||||
):
|
||||
"""
|
||||
Health check for realtime API - tries connection to the realtime API websocket
|
||||
|
|
@ -550,14 +551,17 @@ async def _realtime_health_check(
|
|||
elif custom_llm_provider == "xai":
|
||||
url = xai_realtime._construct_url(api_base=api_base or "https://api.x.ai/v1", query_params={"model": model})
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
vertex_location = litellm.vertex_location or get_secret_str("VERTEXAI_LOCATION")
|
||||
resolved_location = vertex_llm_base.get_vertex_region(vertex_region=vertex_location, model=model)
|
||||
vertex_model_params = model_params or {}
|
||||
resolved_location = vertex_llm_base.get_vertex_region(
|
||||
vertex_region=VertexBase.safe_get_vertex_ai_location(vertex_model_params),
|
||||
model=model,
|
||||
)
|
||||
(
|
||||
access_token,
|
||||
resolved_project,
|
||||
) = await vertex_llm_base._ensure_access_token_async(
|
||||
credentials=None,
|
||||
project_id=litellm.vertex_project or get_secret_str("VERTEXAI_PROJECT"),
|
||||
credentials=VertexBase.safe_get_vertex_ai_credentials(vertex_model_params),
|
||||
project_id=VertexBase.safe_get_vertex_ai_project(vertex_model_params),
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
vertex_realtime_config = VertexAIRealtimeConfig(
|
||||
|
|
|
|||
|
|
@ -277,3 +277,70 @@ async def test_batch_health_check_falls_back_to_acompletion_for_unsupported():
|
|||
)
|
||||
mock_alist.assert_not_called()
|
||||
mock_acompletion.assert_called_once_with(**model_params)
|
||||
|
||||
|
||||
class _FakeWebsocketConnect:
|
||||
def __init__(self, calls, url, **kwargs):
|
||||
calls.append({"url": url, **kwargs})
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realtime_health_check_uses_model_level_vertex_params():
|
||||
"""Regression test: realtime health checks must resolve vertex_credentials,
|
||||
vertex_project, and vertex_location from the model row's params instead of
|
||||
falling back to process-global VERTEXAI_* settings."""
|
||||
import litellm
|
||||
from litellm.realtime_api import main as realtime_main
|
||||
|
||||
fake_vertex_base = MagicMock()
|
||||
fake_vertex_base.get_vertex_region = MagicMock(return_value="us-central1")
|
||||
fake_vertex_base._ensure_access_token_async = AsyncMock(
|
||||
return_value=("model-level-token", "model-level-project")
|
||||
)
|
||||
connect_calls = []
|
||||
|
||||
with (
|
||||
patch.object(realtime_main, "vertex_llm_base", fake_vertex_base),
|
||||
patch(
|
||||
"websockets.connect",
|
||||
lambda url, **kwargs: _FakeWebsocketConnect(connect_calls, url, **kwargs),
|
||||
),
|
||||
patch.object(
|
||||
HealthCheckHelpers,
|
||||
"_update_model_params_with_health_check_tracking_information",
|
||||
staticmethod(lambda model_params: model_params),
|
||||
),
|
||||
):
|
||||
result = await litellm.ahealth_check(
|
||||
model_params={
|
||||
"model": "vertex_ai/gemini-live-2.5-flash-native-audio",
|
||||
"vertex_credentials": '{"type":"service_account"}',
|
||||
"vertex_project": "model-level-project",
|
||||
"vertex_location": "us-central1",
|
||||
},
|
||||
mode="realtime",
|
||||
)
|
||||
|
||||
assert result == {}
|
||||
fake_vertex_base.get_vertex_region.assert_called_once_with(
|
||||
vertex_region="us-central1", model="gemini-live-2.5-flash-native-audio"
|
||||
)
|
||||
fake_vertex_base._ensure_access_token_async.assert_called_once_with(
|
||||
credentials='{"type":"service_account"}',
|
||||
project_id="model-level-project",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert connect_calls[0]["url"] == (
|
||||
"wss://us-central1-aiplatform.googleapis.com/ws/"
|
||||
"google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
|
||||
)
|
||||
assert connect_calls[0]["additional_headers"] == {
|
||||
"Authorization": "Bearer model-level-token",
|
||||
"x-goog-user-project": "model-level-project",
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue