mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
791358912c
commit
8333cfd2f6
2 changed files with 107 additions and 0 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue