fix: evict completed background refresh tasks from _background_refresh_tasks

Completed asyncio.Task objects were never removed from
_background_refresh_tasks. In long-running proxies with many distinct
credential keys the dict grows indefinitely, retaining references to
finished tasks and their results.

Fix:
- Pop the existing (done) entry before creating a replacement task.
- Attach a done_callback to each new task that removes its entry from
  the dict once the task finishes (success or failure).

Tests:
- test_background_refresh_task_removed_after_completion: verifies the
  done-callback cleans up a single entry after the task completes.
- test_background_refresh_tasks_no_accumulation_across_many_keys:
  drives 20 distinct credential keys and confirms the dict is empty
  after all background refreshes finish.

Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-05-04 13:28:08 +00:00
parent 791358912c
commit 8333cfd2f6
No known key found for this signature in database
2 changed files with 107 additions and 0 deletions

View file

@ -900,6 +900,9 @@ class VertexBase:
# 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,
@ -907,6 +910,14 @@ class VertexBase:
credential_project_id,
)
)
# Clean up the entry automatically when the task finishes so
# that long-running proxies with many credential keys do not
# accumulate stale references.
task.add_done_callback(
lambda t, key=credential_cache_key: self._background_refresh_tasks.pop(
key, None
)
)
self._background_refresh_tasks[credential_cache_key] = task
return current_token, resolved_project

View file

@ -1625,6 +1625,102 @@ class TestVertexBase:
assert not mock_refresh.called, "Fresh token should not trigger refresh"
assert token == "fresh-token"
@pytest.mark.asyncio
async def test_background_refresh_task_removed_after_completion(self):
"""Completed background-refresh tasks must be evicted from
_background_refresh_tasks so the dict does not grow unboundedly."""
import asyncio
from google.auth.credentials import TokenState
vertex_base = VertexBase()
mock_creds = MagicMock()
mock_creds.token = "near-expiry-token"
mock_creds.token_state = TokenState.STALE
mock_creds.project_id = "project-1"
mock_creds.quota_project_id = "project-1"
credentials = {"type": "service_account", "project_id": "project-1"}
with (
patch.object(
vertex_base, "load_auth", return_value=(mock_creds, "project-1")
),
patch.object(vertex_base, "refresh_auth") as mock_refresh,
):
def mock_refresh_impl(creds):
creds.token = "refreshed-token"
creds.token_state = TokenState.FRESH
mock_refresh.side_effect = mock_refresh_impl
await vertex_base._ensure_access_token_async(
credentials=credentials,
project_id="project-1",
custom_llm_provider="vertex_ai",
)
# Allow the background task to complete.
await asyncio.sleep(0.1)
# After completion the entry should have been removed by the done-callback.
assert len(vertex_base._background_refresh_tasks) == 0, (
"Completed background refresh task was not removed from "
"_background_refresh_tasks"
)
@pytest.mark.asyncio
async def test_background_refresh_tasks_no_accumulation_across_many_keys(self):
"""With many distinct credential keys the dict must not hold completed tasks."""
import asyncio
import json as _json
from google.auth.credentials import TokenState
vertex_base = VertexBase()
num_keys = 20
for i in range(num_keys):
mock_creds = MagicMock()
mock_creds.token = f"token-{i}"
mock_creds.token_state = TokenState.STALE
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") as mock_refresh,
):
def mock_refresh_impl(creds, idx=i):
creds.token = f"refreshed-{idx}"
creds.token_state = TokenState.FRESH
mock_refresh.side_effect = mock_refresh_impl
await vertex_base._ensure_access_token_async(
credentials=credentials,
project_id=f"project-{i}",
custom_llm_provider="vertex_ai",
)
# Let all background tasks finish.
await asyncio.sleep(0.1)
assert len(vertex_base._background_refresh_tasks) == 0, (
f"Expected 0 tasks after all refreshes completed, "
f"found {len(vertex_base._background_refresh_tasks)}"
)
@pytest.mark.asyncio
async def test_fast_path_no_lock(self):
"""Cached fresh credentials should return without acquiring the lock."""