diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 5002f77abf0..350543f6fa8 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -54,6 +54,9 @@ class VertexBase: # Uses a regular dict (not WeakValueDictionary) so the lock identity is # stable across concurrent callers — a weak reference can be GC'd # between two coroutines arriving at the lock, breaking single-flight. + # Entries are explicitly pruned by _maybe_prune_async_refresh_lock once + # no coroutine holds or is waiting on the lock, so the dict stays + # bounded even in long-running high-cardinality deployments. 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] = {} @@ -369,6 +372,28 @@ class VertexBase: credential_cache_key, asyncio.Lock() ) + def _maybe_prune_async_refresh_lock( + self, credential_cache_key: tuple, lock: asyncio.Lock + ) -> None: + """Drop ``lock`` from ``_async_refresh_locks`` if no coroutine is using it. + + Must be called only after the caller has released ``lock`` (i.e. once + the surrounding ``async with`` has exited). asyncio is cooperative, so + the check-then-pop sequence below runs atomically with respect to other + coroutines — a waiter that arrives later will simply create a fresh + lock under the same key, which is fine because the previous batch is + already done. + """ + if lock.locked(): + return + waiters = getattr(lock, "_waiters", None) + if waiters: + for fut in waiters: + if not fut.cancelled(): + return + if self._async_refresh_locks.get(credential_cache_key) is lock: + self._async_refresh_locks.pop(credential_cache_key, None) + def _try_get_cached_token( self, credential_cache_key: tuple, @@ -475,6 +500,35 @@ class VertexBase: exc_info=True, ) + def _schedule_background_refresh( + self, + credentials: Any, + credential_cache_key: tuple, + credential_project_id: Optional[str], + ) -> None: + """Kick off a single background refresh for ``credential_cache_key``. + + Skips scheduling if a refresh is already in flight. The done-callback + guards against removing a newer task that has replaced this one in the + tracking dict (done_callbacks are scheduled via ``call_soon``). + """ + existing = self._background_refresh_tasks.get(credential_cache_key) + if existing is not None and not existing.done(): + return + self._background_refresh_tasks.pop(credential_cache_key, None) + task = asyncio.create_task( + self._background_refresh_credentials( + credentials, credential_cache_key, credential_project_id + ) + ) + + def _drop_background_refresh_task(_fut: asyncio.Future[Any]) -> None: + if self._background_refresh_tasks.get(credential_cache_key) is _fut: + self._background_refresh_tasks.pop(credential_cache_key, None) + + task.add_done_callback(_drop_background_refresh_task) + self._background_refresh_tasks[credential_cache_key] = task + def _ensure_access_token( self, credentials: Optional[VERTEX_CREDENTIALS_TYPES], @@ -914,120 +968,99 @@ class VertexBase: # === 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 + try: + 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 - ) + _credentials, credential_project_id = self._unpack_cached_credentials( + 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, + # 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 + ) ) - # 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) + # 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, + ) - 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 — block on a full - # refresh using the same path as INVALID credentials. - token_state = TokenState.INVALID - 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(): - # Remove the completed entry before creating a new task so - # that done tasks do not accumulate in the dict indefinitely. - self._background_refresh_tasks.pop(credential_cache_key, None) - task = asyncio.create_task( - self._background_refresh_credentials( - _credentials, - credential_cache_key, - 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 — block on a full + # refresh using the same path as INVALID credentials. + token_state = TokenState.INVALID + else: + self._schedule_background_refresh( + _credentials, + credential_cache_key, + credential_project_id, + ) + 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" ) + return await self._handle_reauthentication_async( + 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__ ) - - # Clean up the entry automatically when the task finishes so - # that long-running proxies with many credential keys do not - # accumulate stale references. Guard with an identity check - # so a stale callback can't remove a newer task that already - # replaced this one in the dict (done_callbacks are scheduled - # via call_soon, so another coroutine may have stored a fresh - # task for the same key before this callback fires). - def _drop_background_refresh_task( - _fut: asyncio.Future[Any], - ) -> None: - if ( - self._background_refresh_tasks.get(credential_cache_key) - is _fut - ): - self._background_refresh_tasks.pop( - credential_cache_key, None - ) - - task.add_done_callback(_drop_background_refresh_task) - 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" - ) - return await self._handle_reauthentication_async( - credentials=credentials, - project_id=project_id, - credential_cache_key=credential_cache_key, - error=e, - ) - raise + if project_id is None: + raise ValueError("Could not resolve project_id") - # 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 + return _credentials.token, project_id + finally: + # Drop the lock from the registry once we're done with it. Safe to + # call here because asyncio.Lock.__aexit__ releases the lock before + # control reaches finally, and asyncio is cooperative so the prune + # check-and-pop is atomic with respect to other coroutines. + self._maybe_prune_async_refresh_lock(credential_cache_key, lock) async def _ensure_access_token_async( self, diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py index 4a1cbb67616..52f0babf16c 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py @@ -1792,6 +1792,79 @@ class TestVertexBase: "receive the same Lock instance to preserve single-flight." ) + @pytest.mark.asyncio + async def test_async_refresh_lock_pruned_after_release(self): + """get_access_token_async must drop the per-key Lock from the registry + once no coroutine is using it, so the dict stays bounded in + high-cardinality deployments. Without this, every distinct credential + leaks a Lock object for the lifetime of the process.""" + from google.auth.credentials import TokenState + + vertex_base = VertexBase() + + for i in range(10): + mock_creds = MagicMock() + mock_creds.token = f"refreshed-{i}" + mock_creds.token_state = TokenState.FRESH + mock_creds.project_id = f"project-{i}" + mock_creds.quota_project_id = f"project-{i}" + + credentials = {"type": "service_account", "project_id": f"project-{i}"} + + with ( + patch.object( + vertex_base, + "load_auth", + return_value=(mock_creds, f"project-{i}"), + ), + patch.object(vertex_base, "refresh_auth"), + ): + await vertex_base._ensure_access_token_async( + credentials=credentials, + project_id=f"project-{i}", + custom_llm_provider="vertex_ai", + ) + + assert len(vertex_base._async_refresh_locks) == 0, ( + "expected per-key locks to be pruned once no coroutine holds or " + f"waits on them; found {len(vertex_base._async_refresh_locks)}" + ) + + @pytest.mark.asyncio + async def test_async_refresh_lock_kept_while_waiter_pending(self): + """The prune must not run while another coroutine is still waiting on + the lock — otherwise the waiter ends up on a lock that's been replaced + in the registry and single-flight breaks.""" + vertex_base = VertexBase() + key = ("creds", "project-1") + lock = vertex_base._get_async_refresh_lock(key) + + async def hold_then_release(release: asyncio.Event): + async with lock: + await release.wait() + + release_holder = asyncio.Event() + holder = asyncio.create_task(hold_then_release(release_holder)) + await asyncio.sleep(0) # let holder grab the lock + + async def queue_for_lock(): + async with lock: + pass + + waiter = asyncio.create_task(queue_for_lock()) + await asyncio.sleep(0) # let waiter enter lock._waiters + + # Prune would normally only run after `async with` exit; call it + # directly while the holder is still active and the waiter is queued. + vertex_base._maybe_prune_async_refresh_lock(key, lock) + assert ( + vertex_base._async_refresh_locks.get(key) is lock + ), "lock with active holder/waiter must not be pruned" + + release_holder.set() + await holder + await waiter + @pytest.mark.asyncio async def test_fast_path_no_lock(self): """Cached fresh credentials should return without acquiring the lock."""