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:
Mateo Wang 2026-07-08 17:30:46 -07:00 • committed by GitHub
parent 4e6ec995e7
commit 5aa5b8ef32
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 76 additions and 4 deletions

View file

@ -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,

View file

@ -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(

View file

@ -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",
}