mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
5862f45a6c
commit
bf23dc5775
2 changed files with 210 additions and 104 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue