mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
refactor(realtime): inject the vertex access token resolver
Take the resolver and its timeout as parameters of the bounded helper and bind the vertex one once at module level, so the timeout tests drive an injected fake instead of patching a shared singleton.
This commit is contained in:
parent
9271133beb
commit
58844d3bda
3 changed files with 86 additions and 41 deletions
|
|
@ -15,7 +15,7 @@ from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
|||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.xai.common_utils import XAIModelInfo
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES
|
||||
from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES, VertexAccessTokenResolver
|
||||
from litellm.types.realtime import (
|
||||
RealtimeClientSecretRequest,
|
||||
RealtimeExpiresAfter,
|
||||
|
|
@ -43,6 +43,7 @@ openai_realtime: Final = OpenAIRealtime()
|
|||
bedrock_realtime: Final = BedrockRealtime()
|
||||
xai_realtime: Final = XAIRealtime()
|
||||
vertex_llm_base: Final = VertexBase()
|
||||
vertex_access_token_resolver: Final[VertexAccessTokenResolver] = vertex_llm_base._ensure_access_token_async
|
||||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
|
||||
|
||||
|
|
@ -290,20 +291,22 @@ async def arealtime_calls(
|
|||
async def _resolve_vertex_access_token_bounded(
|
||||
credentials: VERTEX_CREDENTIALS_TYPES | None,
|
||||
project_id: str | None,
|
||||
resolver: VertexAccessTokenResolver,
|
||||
timeout_seconds: float,
|
||||
) -> tuple[str, str]:
|
||||
try:
|
||||
return await asyncio.wait_for(
|
||||
vertex_llm_base._ensure_access_token_async(
|
||||
resolver(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
custom_llm_provider="vertex_ai",
|
||||
),
|
||||
timeout=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
||||
timeout=timeout_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ValueError(
|
||||
"Vertex AI realtime: timed out fetching Google OAuth access token after "
|
||||
f"{REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS}s; check network egress from the proxy "
|
||||
f"{timeout_seconds}s; check network egress from the proxy "
|
||||
"to the OAuth token endpoint (oauth2.googleapis.com)"
|
||||
) from e
|
||||
|
||||
|
|
@ -508,6 +511,8 @@ async def _arealtime(
|
|||
) = await _resolve_vertex_access_token_bounded(
|
||||
credentials=vertex_credentials,
|
||||
project_id=vertex_project,
|
||||
resolver=vertex_access_token_resolver,
|
||||
timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
vertex_realtime_config: Final = VertexAIRealtimeConfig(
|
||||
|
|
@ -588,6 +593,8 @@ async def _realtime_health_check(
|
|||
) = await _resolve_vertex_access_token_bounded(
|
||||
credentials=VertexBase.safe_get_vertex_ai_credentials(vertex_model_params),
|
||||
project_id=VertexBase.safe_get_vertex_ai_project(vertex_model_params),
|
||||
resolver=vertex_access_token_resolver,
|
||||
timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
||||
)
|
||||
vertex_realtime_config: Final = VertexAIRealtimeConfig(
|
||||
access_token=access_token,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from enum import Enum
|
||||
from typing import Any, Final, Literal
|
||||
from typing import Any, Final, Literal, Protocol
|
||||
|
||||
from typing_extensions import (
|
||||
Required,
|
||||
|
|
@ -747,6 +747,17 @@ class VertexVideoGenerationResponse(TypedDict, total=False):
|
|||
VERTEX_CREDENTIALS_TYPES = str | dict[str, str]
|
||||
|
||||
|
||||
class VertexAccessTokenResolver(Protocol):
|
||||
"""Resolves a Google OAuth access token and the project id it belongs to."""
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
credentials: VERTEX_CREDENTIALS_TYPES | None,
|
||||
project_id: str | None,
|
||||
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
|
||||
) -> tuple[str, str]: ...
|
||||
|
||||
|
||||
class VertexPartnerProvider(str, Enum):
|
||||
mistralai = "mistralai"
|
||||
llama = "llama"
|
||||
|
|
|
|||
|
|
@ -93,12 +93,70 @@ def test_client_secret_session_model_takes_priority_over_top_level(monkeypatch):
|
|||
assert captured["request_data"]["session"]["model"] == "gpt-realtime-session"
|
||||
|
||||
|
||||
async def _hanging_resolver(credentials, project_id, custom_llm_provider) -> tuple[str, str]:
|
||||
await asyncio.sleep(30)
|
||||
return "", ""
|
||||
|
||||
|
||||
async def _thread_offloaded_hanging_resolver(credentials, project_id, custom_llm_provider) -> tuple[str, str]:
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
|
||||
await asyncify(time.sleep)(30)
|
||||
return "", ""
|
||||
|
||||
|
||||
async def _instant_resolver(credentials, project_id, custom_llm_provider) -> tuple[str, str]:
|
||||
return "token-abc", "resolved-project"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arealtime_vertex_hung_credential_resolution_raises_promptly(monkeypatch):
|
||||
async def test_vertex_credential_resolution_returns_the_resolved_token_and_project():
|
||||
assert await realtime_main._resolve_vertex_access_token_bounded(
|
||||
credentials="fake-credentials",
|
||||
project_id="fake-project",
|
||||
resolver=_instant_resolver,
|
||||
timeout_seconds=5,
|
||||
) == ("token-abc", "resolved-project")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_credential_resolution_times_out_instead_of_hanging():
|
||||
"""Regression for the realtime accept-then-silence hang: a stalled Google
|
||||
OAuth token refresh used to block _arealtime's vertex branch unbounded
|
||||
(minutes of zero frames for the client). It must instead raise a clear,
|
||||
prompt error naming the credential-resolution timeout."""
|
||||
OAuth token refresh used to block the vertex branch unbounded (minutes of
|
||||
zero frames for the client). It must raise promptly and name the timeout."""
|
||||
start = time.monotonic()
|
||||
with pytest.raises(ValueError, match="timed out fetching Google OAuth access token"):
|
||||
await realtime_main._resolve_vertex_access_token_bounded(
|
||||
credentials="fake-credentials",
|
||||
project_id="fake-project",
|
||||
resolver=_hanging_resolver,
|
||||
timeout_seconds=0.05,
|
||||
)
|
||||
assert time.monotonic() - start < 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_credential_resolution_bounds_a_thread_offloaded_refresh():
|
||||
"""The real stall is a blocking google-auth refresh that runs in a worker
|
||||
thread via asyncify, not a plain awaitable sleep. A timeout that only bounds
|
||||
cancellable awaits would leave that shape hanging, so bound the shape the
|
||||
proxy actually runs."""
|
||||
start = time.monotonic()
|
||||
with pytest.raises(ValueError, match="timed out fetching Google OAuth access token"):
|
||||
await realtime_main._resolve_vertex_access_token_bounded(
|
||||
credentials="fake-credentials",
|
||||
project_id="fake-project",
|
||||
resolver=_thread_offloaded_hanging_resolver,
|
||||
timeout_seconds=0.05,
|
||||
)
|
||||
assert time.monotonic() - start < 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arealtime_vertex_branch_resolves_credentials_under_a_bound(monkeypatch):
|
||||
"""The wiring half of the regression: the vertex branch of _arealtime must
|
||||
go through the bounded resolver, so a hung token refresh surfaces as a
|
||||
prompt error there rather than as an accepted-then-silent websocket."""
|
||||
|
||||
async def hanging_token_refresh(**kwargs):
|
||||
await asyncio.sleep(30)
|
||||
|
|
@ -107,38 +165,7 @@ async def test_arealtime_vertex_hung_credential_resolution_raises_promptly(monke
|
|||
return model, "vertex_ai", None, api_base
|
||||
|
||||
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
|
||||
monkeypatch.setattr(realtime_main.vertex_llm_base, "_ensure_access_token_async", hanging_token_refresh)
|
||||
monkeypatch.setattr(realtime_main, "REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", 0.05)
|
||||
|
||||
start = time.monotonic()
|
||||
with pytest.raises(ValueError, match="timed out fetching Google OAuth access token"):
|
||||
await realtime_main._arealtime.__wrapped__(
|
||||
model="gemini-live-2.5-flash",
|
||||
websocket=MagicMock(),
|
||||
litellm_logging_obj=FakeLogging(),
|
||||
vertex_credentials="fake-credentials",
|
||||
vertex_project="fake-project",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
assert time.monotonic() - start < 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arealtime_vertex_credential_timeout_survives_thread_offloaded_refresh(monkeypatch):
|
||||
"""The real stall is a blocking google-auth refresh that runs in a worker
|
||||
thread via asyncify, not a plain awaitable sleep. A timeout that only bounds
|
||||
cancellable awaits would leave that shape hanging, so bound the shape the
|
||||
proxy actually runs."""
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
|
||||
async def thread_offloaded_hanging_refresh(**kwargs):
|
||||
return await asyncify(time.sleep)(30)
|
||||
|
||||
def mock_get_llm_provider(model, api_base, api_key):
|
||||
return model, "vertex_ai", None, api_base
|
||||
|
||||
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
|
||||
monkeypatch.setattr(realtime_main.vertex_llm_base, "_ensure_access_token_async", thread_offloaded_hanging_refresh)
|
||||
monkeypatch.setattr(realtime_main, "vertex_access_token_resolver", hanging_token_refresh)
|
||||
monkeypatch.setattr(realtime_main, "REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", 0.05)
|
||||
|
||||
start = time.monotonic()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue