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:
mateo-berri 2026-05-20 20:39:08 +00:00
parent 2ff03f3842
commit 0570b16b75
No known key found for this signature in database
2 changed files with 63 additions and 47 deletions

View file

@ -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,

View file

@ -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,