fix(greptile): prune vertex async refresh lock dict after release

Address greptile's open thread on _async_refresh_locks growing
unboundedly in high-cardinality deployments.

- Add _maybe_prune_async_refresh_lock: drops the per-key Lock from
  the registry once no coroutine holds it and no coroutine is queued
  in lock._waiters. The check-then-pop sequence is safe under
  asyncio's cooperative scheduler — a waiter that arrives after the
  pop simply creates a fresh lock under the same key, which is fine
  because the previous batch is already done.
- Wrap the slow-path async with lock in a try/finally so the prune
  runs on every exit (return, exception, reauth retry).
- Extract the existing background-refresh task scheduling into
  _schedule_background_refresh so get_access_token_async stays under
  ruff's PLR0915 ("Too many statements") limit. No behaviour change.
- Regression tests cover both pruning after release (the dict
  shrinks back to zero after each call) and the safeguard that
  keeps the lock alive while a waiter is still queued.
This commit is contained in:
Claude 2026-05-20 19:11:02 +00:00 • committed by Claude
parent 5862f45a6c
commit bf23dc5775
No known key found for this signature in database
2 changed files with 210 additions and 104 deletions

View file

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

View file

@ -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."""