From 0570b16b75a8d6012c2880f31cd82dfdaa9211cd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 20 May 2026 20:39:08 +0000 Subject: [PATCH] fix(greptile): refcount vertex async refresh lock pruning Replace the asyncio.Lock._waiters inspection in _maybe_prune_async_refresh_lock with an explicit refcount so the entry is pruned exactly when no coroutine is holding or waiting on the lock, without depending on any private asyncio internals. --- litellm/llms/vertex_ai/vertex_llm_base.py | 51 ++++++++-------- .../llms/vertex_ai/test_vertex_llm_base.py | 59 +++++++++++-------- 2 files changed, 63 insertions(+), 47 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 16533125b6b..f21ad6ac071 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -54,10 +54,12 @@ 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. + # An explicit refcount tracks the number of coroutines currently using + # each lock; the entry is pruned when the count reaches zero, so the + # dict stays bounded even in long-running high-cardinality deployments + # without depending on any private asyncio internals. self._async_refresh_locks: Dict[tuple, asyncio.Lock] = {} + self._async_refresh_lock_refcounts: Dict[tuple, int] = {} # 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. @@ -366,31 +368,36 @@ 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( + def _acquire_async_refresh_lock(self, credential_cache_key: tuple) -> asyncio.Lock: + """Increment the refcount and return the lock for ``credential_cache_key``. + + Every call must be paired with ``_release_async_refresh_lock`` once the + caller is done with the lock so the entry can be pruned when no other + coroutine is holding or waiting on it. + """ + lock = self._async_refresh_locks.setdefault( credential_cache_key, asyncio.Lock() ) + self._async_refresh_lock_refcounts[credential_cache_key] = ( + self._async_refresh_lock_refcounts.get(credential_cache_key, 0) + 1 + ) + return lock - def _maybe_prune_async_refresh_lock( + def _release_async_refresh_lock( self, credential_cache_key: tuple, lock: asyncio.Lock ) -> None: - """Drop ``lock`` from ``_async_refresh_locks`` if no coroutine is using it. + """Decrement the refcount and drop the lock entry when it reaches zero. 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. + the decrement-then-pop sequence below runs atomically with respect to + other coroutines. """ - if lock.locked(): + remaining = self._async_refresh_lock_refcounts.get(credential_cache_key, 0) - 1 + if remaining > 0: + self._async_refresh_lock_refcounts[credential_cache_key] = remaining return - waiters = getattr(lock, "_waiters", None) - if waiters: - for fut in waiters: - if not fut.cancelled(): - return + self._async_refresh_lock_refcounts.pop(credential_cache_key, None) if self._async_refresh_locks.get(credential_cache_key) is lock: self._async_refresh_locks.pop(credential_cache_key, None) @@ -988,7 +995,7 @@ class VertexBase: return cached # === SLOW PATH (per-key lock) === - lock = self._get_async_refresh_lock(credential_cache_key) + lock = self._acquire_async_refresh_lock(credential_cache_key) try: async with lock: # Double-check after acquiring lock — another coroutine may have refreshed. @@ -1089,11 +1096,7 @@ class VertexBase: 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) + self._release_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 52f0babf16c..2cf97081806 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 @@ -1784,13 +1784,19 @@ class TestVertexBase: vertex_base = VertexBase() key = ("creds", "project-1") - lock_a = vertex_base._get_async_refresh_lock(key) - async with lock_a: - lock_b = vertex_base._get_async_refresh_lock(key) - assert lock_a is lock_b, ( - "While a coroutine still holds the lock, concurrent callers must " - "receive the same Lock instance to preserve single-flight." - ) + lock_a = vertex_base._acquire_async_refresh_lock(key) + try: + async with lock_a: + lock_b = vertex_base._acquire_async_refresh_lock(key) + try: + assert lock_a is lock_b, ( + "While a coroutine still holds the lock, concurrent callers must " + "receive the same Lock instance to preserve single-flight." + ) + finally: + vertex_base._release_async_refresh_lock(key, lock_b) + finally: + vertex_base._release_async_refresh_lock(key, lock_a) @pytest.mark.asyncio async def test_async_refresh_lock_pruned_after_release(self): @@ -1829,6 +1835,7 @@ class TestVertexBase: "expected per-key locks to be pruned once no coroutine holds or " f"waits on them; found {len(vertex_base._async_refresh_locks)}" ) + assert len(vertex_base._async_refresh_lock_refcounts) == 0 @pytest.mark.asyncio async def test_async_refresh_lock_kept_while_waiter_pending(self): @@ -1837,34 +1844,40 @@ class TestVertexBase: 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() + holder_lock = vertex_base._acquire_async_refresh_lock(key) release_holder = asyncio.Event() - holder = asyncio.create_task(hold_then_release(release_holder)) + + async def hold_then_release(): + async with holder_lock: + await release_holder.wait() + vertex_base._release_async_refresh_lock(key, holder_lock) + + holder = asyncio.create_task(hold_then_release()) await asyncio.sleep(0) # let holder grab the lock async def queue_for_lock(): - async with lock: - pass + waiter_lock = vertex_base._acquire_async_refresh_lock(key) + try: + async with waiter_lock: + pass + finally: + vertex_base._release_async_refresh_lock(key, waiter_lock) waiter = asyncio.create_task(queue_for_lock()) - await asyncio.sleep(0) # let waiter enter lock._waiters + await asyncio.sleep(0) # let waiter queue on the lock - # 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 + vertex_base._async_refresh_locks.get(key) is holder_lock ), "lock with active holder/waiter must not be pruned" release_holder.set() await holder await waiter + assert key not in vertex_base._async_refresh_locks + assert key not in vertex_base._async_refresh_lock_refcounts + @pytest.mark.asyncio async def test_fast_path_no_lock(self): """Cached fresh credentials should return without acquiring the lock.""" @@ -1893,11 +1906,11 @@ class TestVertexBase: "project-1", ) - # Spy on _get_async_refresh_lock to verify it's never called + # Spy on _acquire_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, + "_acquire_async_refresh_lock", + wraps=vertex_base._acquire_async_refresh_lock, ) as mock_get_lock: token, project = await vertex_base._ensure_access_token_async( credentials=credentials,