mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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.
This commit is contained in:
parent
2ff03f3842
commit
0570b16b75
2 changed files with 63 additions and 47 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue