mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(vertex_ai): single-flight credential refresh to prevent thundering herd (#26024)
* fix(vertex_ai): single-flight credential refresh to prevent thundering herd When GCP credentials expire under high concurrency, all requests simultaneously call credentials.refresh() via asyncify, saturating the 40-thread anyio pool and blocking the proxy for 20+ seconds. This adds: - Per-credential asyncio.Lock in get_access_token_async for single-flight refresh (1 coroutine refreshes, others wait on the lock) - Background refresh when token_state is STALE (usable but near expiry), returning the current token immediately with zero added latency - threading.Lock on the sync get_access_token path - Uses google-auth's TokenState enum (FRESH/STALE/INVALID) instead of reimplementing expiry logic Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> * fix: address PR review comments - Use asyncio.create_task() instead of deprecated get_event_loop().create_task() - Track in-flight background refresh tasks to prevent duplicate refreshes when multiple STALE-path callers pass through the lock before the first background task completes - Add token validation in the STALE branch (consistent with FRESH/INVALID) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> * fix: lazy-import TokenState to avoid breaking when google-auth is not installed Also extract helper methods to bring get_access_token_async under the PLR0915 statement limit (50). Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> * chore: apply Black formatting to test file and update uv.lock Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> * fix: remove user-provided project_id from log messages (CodeQL log injection) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> * fix: avoid leaking token value in error message, log type instead Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> * chore: restore uv.lock to match litellm_oss_branch Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> * fix: remove project_id from remaining log message (CodeQL log injection) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> * fix: remove remaining project_id from log and error messages Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
26fcbc93e5
commit
4583310313
2 changed files with 453 additions and 29 deletions
|
|
@ -4,8 +4,10 @@ Base Vertex, Google AI Studio LLM Class
|
|||
Handles Authentication and generating request urls for Vertex AI and Google AI Studio
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
|
||||
import litellm
|
||||
|
|
@ -30,6 +32,7 @@ GOOGLE_IMPORT_ERROR_MESSAGE = (
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from google.auth.credentials import Credentials as GoogleCredentialsObject
|
||||
from google.auth.credentials import TokenState
|
||||
else:
|
||||
GoogleCredentialsObject = Any
|
||||
|
||||
|
|
@ -46,6 +49,13 @@ class VertexBase:
|
|||
] = {}
|
||||
self.project_id: Optional[str] = None
|
||||
self.async_handler: Optional[AsyncHTTPHandler] = None
|
||||
# Per-credential-key asyncio.Lock for single-flight async refresh.
|
||||
# Prevents thundering herd when token expires under high concurrency.
|
||||
self._async_refresh_locks: Dict[tuple, asyncio.Lock] = {}
|
||||
# Tracks in-flight background refresh tasks to avoid duplicate refreshes.
|
||||
self._background_refresh_tasks: Dict[tuple, asyncio.Task] = {}
|
||||
# Protects the sync get_access_token refresh path.
|
||||
self._sync_refresh_lock = threading.Lock()
|
||||
|
||||
def get_vertex_region(self, vertex_region: Optional[str], model: str) -> str:
|
||||
import litellm
|
||||
|
|
@ -77,7 +87,9 @@ class VertexBase:
|
|||
return vertex_region or "us-central1"
|
||||
|
||||
def load_auth(
|
||||
self, credentials: Optional[VERTEX_CREDENTIALS_TYPES], project_id: Optional[str]
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
project_id: Optional[str],
|
||||
) -> Tuple[Any, str]:
|
||||
if credentials is not None:
|
||||
if isinstance(credentials, str):
|
||||
|
|
@ -345,6 +357,118 @@ class VertexBase:
|
|||
|
||||
credentials.refresh(Request())
|
||||
|
||||
def _get_async_refresh_lock(self, credential_cache_key: tuple) -> asyncio.Lock:
|
||||
"""Get or create an asyncio.Lock for the given credential cache key."""
|
||||
return self._async_refresh_locks.setdefault(
|
||||
credential_cache_key, asyncio.Lock()
|
||||
)
|
||||
|
||||
def _try_get_cached_token(
|
||||
self,
|
||||
credential_cache_key: tuple,
|
||||
project_id: Optional[str],
|
||||
) -> Optional[Tuple[str, str]]:
|
||||
"""
|
||||
Look up cached credentials and return (token, project_id) if the token
|
||||
is FRESH. Returns None if not cached or not fresh.
|
||||
"""
|
||||
from google.auth.credentials import TokenState
|
||||
|
||||
creds, cached_project_id = self._unpack_cached_credentials(credential_cache_key)
|
||||
if (
|
||||
creds is not None
|
||||
and self._get_token_state(creds) == TokenState.FRESH
|
||||
and creds.token is not None
|
||||
and isinstance(creds.token, str)
|
||||
):
|
||||
resolved_project = project_id or cached_project_id
|
||||
if resolved_project:
|
||||
return creds.token, resolved_project
|
||||
return None
|
||||
|
||||
def _unpack_cached_credentials(
|
||||
self, credential_cache_key: tuple
|
||||
) -> Tuple[Any, Optional[str]]:
|
||||
"""
|
||||
Return (credentials, project_id) from the cache, or (None, None) if
|
||||
not cached. Handles both tuple and legacy cache formats.
|
||||
"""
|
||||
if credential_cache_key not in self._credentials_project_mapping:
|
||||
return None, None
|
||||
cached_entry = self._credentials_project_mapping[credential_cache_key]
|
||||
if isinstance(cached_entry, tuple):
|
||||
return cached_entry
|
||||
return cached_entry, cached_entry.quota_project_id or getattr(
|
||||
cached_entry, "project_id", None
|
||||
)
|
||||
|
||||
def _get_token_state(self, credentials: Any) -> "TokenState":
|
||||
"""
|
||||
Return the token state using google-auth's TokenState enum.
|
||||
|
||||
Falls back to expired/valid checks if token_state is unavailable
|
||||
(e.g. older google-auth versions or mock objects in tests).
|
||||
"""
|
||||
from google.auth.credentials import TokenState as _TokenState
|
||||
|
||||
token_state = getattr(credentials, "token_state", None)
|
||||
if isinstance(token_state, _TokenState):
|
||||
return token_state
|
||||
# Fallback for credentials without a real token_state (e.g. mocks)
|
||||
if getattr(credentials, "expired", True):
|
||||
return _TokenState.INVALID
|
||||
if getattr(credentials, "valid", False):
|
||||
return _TokenState.FRESH
|
||||
return _TokenState.INVALID
|
||||
|
||||
async def _load_and_cache_credentials(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
project_id: Optional[str],
|
||||
credential_cache_key: tuple,
|
||||
) -> Tuple[Any, Optional[str]]:
|
||||
"""Load credentials via load_auth (in thread) and cache the result."""
|
||||
try:
|
||||
_credentials, credential_project_id = await asyncify(self.load_auth)(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Failed to load vertex credentials: %s", str(e))
|
||||
raise
|
||||
if _credentials is None:
|
||||
raise ValueError("Could not resolve credentials")
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
return _credentials, credential_project_id
|
||||
|
||||
async def _background_refresh_credentials(
|
||||
self,
|
||||
credentials: Any,
|
||||
credential_cache_key: tuple,
|
||||
credential_project_id: Optional[str],
|
||||
) -> None:
|
||||
"""
|
||||
Refresh credentials in the background without blocking the calling request.
|
||||
|
||||
Called when the token is still valid but nearing expiry (proactive refresh).
|
||||
Errors are logged but not raised — the current token is still usable.
|
||||
"""
|
||||
try:
|
||||
verbose_logger.debug("Background proactive credential refresh")
|
||||
await asyncify(self.refresh_auth)(credentials)
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
except Exception:
|
||||
verbose_logger.debug(
|
||||
"Background credential refresh failed, will retry on next request",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
def _ensure_access_token(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
|
|
@ -646,7 +770,7 @@ class VertexBase:
|
|||
)
|
||||
|
||||
## VALIDATE CREDENTIALS
|
||||
verbose_logger.debug(f"Validating credentials for project_id: {project_id}")
|
||||
verbose_logger.debug("Validating credentials")
|
||||
if (
|
||||
project_id is None
|
||||
and credential_project_id is not None
|
||||
|
|
@ -666,26 +790,27 @@ class VertexBase:
|
|||
raise ValueError("Credentials are None after loading")
|
||||
|
||||
if _credentials.expired:
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
f"Credentials expired, refreshing for project_id: {project_id}"
|
||||
)
|
||||
self.refresh_auth(_credentials)
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
# if refresh fails, it's possible the user has re-authenticated via `gcloud auth application-default login`
|
||||
# in this case, we should try to reload the credentials by clearing the cache and retrying
|
||||
if "Reauthentication is needed" in str(e) and not _retry_reauth:
|
||||
return self._handle_reauthentication(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
error=e,
|
||||
)
|
||||
raise e
|
||||
with self._sync_refresh_lock:
|
||||
# Double-check after acquiring lock
|
||||
if _credentials.expired:
|
||||
try:
|
||||
verbose_logger.debug("Credentials expired, refreshing")
|
||||
self.refresh_auth(_credentials)
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
# if refresh fails, it's possible the user has re-authenticated via `gcloud auth application-default login`
|
||||
# in this case, we should try to reload the credentials by clearing the cache and retrying
|
||||
if "Reauthentication is needed" in str(e) and not _retry_reauth:
|
||||
return self._handle_reauthentication(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
error=e,
|
||||
)
|
||||
raise e
|
||||
|
||||
## VALIDATION STEP
|
||||
if _credentials.token is None or not isinstance(_credentials.token, str):
|
||||
|
|
@ -700,6 +825,127 @@ class VertexBase:
|
|||
|
||||
return _credentials.token, project_id
|
||||
|
||||
async def get_access_token_async(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
project_id: Optional[str],
|
||||
) -> Tuple[str, str]:
|
||||
"""
|
||||
Async version of get_access_token with single-flight refresh coordination.
|
||||
|
||||
Prevents thundering herd: when credentials expire under high concurrency,
|
||||
only one coroutine refreshes while others wait on the lock. Uses native
|
||||
async refresh for service_account and authorized_user credentials.
|
||||
"""
|
||||
from google.auth.credentials import TokenState
|
||||
|
||||
cache_credentials = (
|
||||
json.dumps(credentials) if isinstance(credentials, dict) else credentials
|
||||
)
|
||||
credential_cache_key = (cache_credentials, project_id)
|
||||
|
||||
# === FAST PATH (no lock) ===
|
||||
# If credentials are FRESH (valid, not near expiry), return immediately.
|
||||
cached = self._try_get_cached_token(credential_cache_key, project_id)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
# === SLOW PATH (per-key lock) ===
|
||||
lock = self._get_async_refresh_lock(credential_cache_key)
|
||||
async with lock:
|
||||
# Double-check after acquiring lock — another coroutine may have refreshed.
|
||||
cached = self._try_get_cached_token(credential_cache_key, project_id)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
_credentials, credential_project_id = self._unpack_cached_credentials(
|
||||
credential_cache_key
|
||||
)
|
||||
|
||||
# Load credentials if not cached
|
||||
if _credentials is None:
|
||||
_credentials, credential_project_id = (
|
||||
await self._load_and_cache_credentials(
|
||||
credentials, project_id, credential_cache_key
|
||||
)
|
||||
)
|
||||
|
||||
# Resolve project_id from credentials if not provided
|
||||
if project_id is None and isinstance(credential_project_id, str):
|
||||
project_id = credential_project_id
|
||||
resolved_cache_key = (cache_credentials, project_id)
|
||||
if resolved_cache_key not in self._credentials_project_mapping:
|
||||
self._credentials_project_mapping[resolved_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
|
||||
# Use google-auth's token_state to decide refresh strategy:
|
||||
# - STALE: token is usable but within REFRESH_THRESHOLD (3:45) of
|
||||
# expiry — return it immediately and refresh in the background.
|
||||
# - INVALID: token is expired or missing — must block on refresh.
|
||||
token_state = self._get_token_state(_credentials)
|
||||
|
||||
if token_state == TokenState.STALE:
|
||||
resolved_project = project_id
|
||||
if resolved_project is None:
|
||||
raise ValueError("Could not resolve project_id")
|
||||
current_token = _credentials.token
|
||||
if current_token is None or not isinstance(current_token, str):
|
||||
# Token is malformed despite STALE state — fall through
|
||||
# to INVALID path which will block on a full refresh.
|
||||
pass
|
||||
else:
|
||||
# Schedule a single background refresh — skip if one is
|
||||
# already in flight for this credential key.
|
||||
existing = self._background_refresh_tasks.get(credential_cache_key)
|
||||
if existing is None or existing.done():
|
||||
task = asyncio.create_task(
|
||||
self._background_refresh_credentials(
|
||||
_credentials,
|
||||
credential_cache_key,
|
||||
credential_project_id,
|
||||
)
|
||||
)
|
||||
self._background_refresh_tasks[credential_cache_key] = task
|
||||
return current_token, resolved_project
|
||||
|
||||
if token_state == TokenState.INVALID:
|
||||
# Token is expired or missing — must block until refresh completes.
|
||||
try:
|
||||
verbose_logger.debug("Credentials expired, refreshing")
|
||||
await asyncify(self.refresh_auth)(_credentials)
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
if "Reauthentication is needed" in str(e):
|
||||
verbose_logger.debug(
|
||||
"Reauthentication needed, clearing cache and retrying"
|
||||
)
|
||||
if credential_cache_key in self._credentials_project_mapping:
|
||||
del self._credentials_project_mapping[credential_cache_key]
|
||||
return await asyncify(self._handle_reauthentication)(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
error=e,
|
||||
)
|
||||
raise
|
||||
|
||||
# Final validation
|
||||
if _credentials.token is None or not isinstance(_credentials.token, str):
|
||||
raise ValueError(
|
||||
"Could not resolve credentials token. Got None or non-string token (type={})".format(
|
||||
type(_credentials.token).__name__
|
||||
)
|
||||
)
|
||||
if project_id is None:
|
||||
raise ValueError("Could not resolve project_id")
|
||||
|
||||
return _credentials.token, project_id
|
||||
|
||||
async def _ensure_access_token_async(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
|
|
@ -714,13 +960,10 @@ class VertexBase:
|
|||
if custom_llm_provider == "gemini":
|
||||
return "", ""
|
||||
else:
|
||||
try:
|
||||
return await asyncify(self.get_access_token)(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
raise e
|
||||
return await self.get_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
)
|
||||
|
||||
def set_headers(
|
||||
self, auth_header: Optional[str], extra_headers: Optional[dict]
|
||||
|
|
|
|||
|
|
@ -1448,3 +1448,184 @@ class TestVertexBase:
|
|||
|
||||
aws_creds = supplier.get_aws_security_credentials(context=None, request=None)
|
||||
assert isinstance(aws_creds, AwsSecurityCredentials)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_flight_refresh(self):
|
||||
"""Under high concurrency, only one coroutine should refresh expired credentials."""
|
||||
import asyncio
|
||||
|
||||
vertex_base = VertexBase()
|
||||
|
||||
mock_creds = MagicMock()
|
||||
mock_creds.token = "expired-token"
|
||||
mock_creds.expired = True
|
||||
mock_creds.expiry = None
|
||||
mock_creds.project_id = "project-1"
|
||||
mock_creds.quota_project_id = "project-1"
|
||||
|
||||
credentials = {"type": "service_account", "project_id": "project-1"}
|
||||
|
||||
refresh_call_count = 0
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
vertex_base, "load_auth", return_value=(mock_creds, "project-1")
|
||||
),
|
||||
patch.object(vertex_base, "refresh_auth") as mock_refresh,
|
||||
):
|
||||
|
||||
async def slow_refresh(creds):
|
||||
nonlocal refresh_call_count
|
||||
refresh_call_count += 1
|
||||
await asyncio.sleep(0.05) # simulate network latency
|
||||
creds.token = "refreshed-token"
|
||||
creds.expired = False
|
||||
|
||||
# refresh_auth is sync, but we need to count calls.
|
||||
# get_access_token_async wraps it with asyncify, so the sync side_effect works.
|
||||
def sync_refresh_impl(creds):
|
||||
nonlocal refresh_call_count
|
||||
refresh_call_count += 1
|
||||
creds.token = "refreshed-token"
|
||||
creds.expired = False
|
||||
|
||||
mock_refresh.side_effect = sync_refresh_impl
|
||||
|
||||
# Launch 50 concurrent requests
|
||||
tasks = [
|
||||
vertex_base._ensure_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id="project-1",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
for _ in range(50)
|
||||
]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
# All should return the refreshed token
|
||||
for token, project in results:
|
||||
assert token == "refreshed-token"
|
||||
assert project == "project-1"
|
||||
|
||||
# refresh_auth should be called exactly once (single-flight)
|
||||
assert (
|
||||
refresh_call_count == 1
|
||||
), f"Expected 1 refresh call, got {refresh_call_count}"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_refresh_when_near_expiry(self):
|
||||
"""When token_state is STALE (within the 3:45 REFRESH_THRESHOLD window),
|
||||
return the current token immediately and refresh in the background —
|
||||
zero added latency."""
|
||||
import asyncio
|
||||
|
||||
from google.auth.credentials import TokenState
|
||||
|
||||
vertex_base = VertexBase()
|
||||
|
||||
# Simulate STALE state: token is usable but near expiry.
|
||||
mock_creds = MagicMock()
|
||||
mock_creds.token = "near-expiry-token"
|
||||
mock_creds.token_state = TokenState.STALE
|
||||
mock_creds.project_id = "project-1"
|
||||
mock_creds.quota_project_id = "project-1"
|
||||
|
||||
credentials = {"type": "service_account", "project_id": "project-1"}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
vertex_base, "load_auth", return_value=(mock_creds, "project-1")
|
||||
),
|
||||
patch.object(vertex_base, "refresh_auth") as mock_refresh,
|
||||
):
|
||||
|
||||
def mock_refresh_impl(creds):
|
||||
creds.token = "refreshed-token"
|
||||
creds.token_state = TokenState.FRESH
|
||||
|
||||
mock_refresh.side_effect = mock_refresh_impl
|
||||
|
||||
token, project = await vertex_base._ensure_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id="project-1",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
# Should return the current (still usable) token immediately
|
||||
assert token == "near-expiry-token"
|
||||
|
||||
# Let the background refresh task run
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
assert mock_refresh.called, "Background refresh should have been triggered"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fresh_token_skips_refresh(self):
|
||||
"""Credentials not marked expired by google-auth should not trigger refresh."""
|
||||
vertex_base = VertexBase()
|
||||
|
||||
mock_creds = MagicMock()
|
||||
mock_creds.token = "fresh-token"
|
||||
mock_creds.expired = False
|
||||
mock_creds.project_id = "project-1"
|
||||
mock_creds.quota_project_id = "project-1"
|
||||
|
||||
credentials = {"type": "service_account", "project_id": "project-1"}
|
||||
cache_key = (json.dumps(credentials), "project-1")
|
||||
vertex_base._credentials_project_mapping[cache_key] = (
|
||||
mock_creds,
|
||||
"project-1",
|
||||
)
|
||||
|
||||
with patch.object(vertex_base, "refresh_auth") as mock_refresh:
|
||||
token, project = await vertex_base._ensure_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id="project-1",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
assert not mock_refresh.called, "Fresh token should not trigger refresh"
|
||||
assert token == "fresh-token"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fast_path_no_lock(self):
|
||||
"""Cached fresh credentials should return without acquiring the lock."""
|
||||
import datetime
|
||||
|
||||
vertex_base = VertexBase()
|
||||
|
||||
try:
|
||||
from google.auth import _helpers as google_auth_helpers
|
||||
|
||||
now = google_auth_helpers.utcnow()
|
||||
except ImportError:
|
||||
now = datetime.datetime.utcnow()
|
||||
|
||||
mock_creds = MagicMock()
|
||||
mock_creds.token = "cached-token"
|
||||
mock_creds.expired = False
|
||||
mock_creds.expiry = now + datetime.timedelta(minutes=30)
|
||||
mock_creds.project_id = "project-1"
|
||||
mock_creds.quota_project_id = "project-1"
|
||||
|
||||
credentials = {"type": "service_account", "project_id": "project-1"}
|
||||
cache_key = (json.dumps(credentials), "project-1")
|
||||
vertex_base._credentials_project_mapping[cache_key] = (
|
||||
mock_creds,
|
||||
"project-1",
|
||||
)
|
||||
|
||||
# Spy on _get_async_refresh_lock to verify it's never called
|
||||
with patch.object(
|
||||
vertex_base,
|
||||
"_get_async_refresh_lock",
|
||||
wraps=vertex_base._get_async_refresh_lock,
|
||||
) as mock_get_lock:
|
||||
token, project = await vertex_base._ensure_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id="project-1",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
assert token == "cached-token"
|
||||
assert not mock_get_lock.called, "Fast path should not acquire lock"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue