mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix: preserve async batch and reauth coordination
This commit is contained in:
parent
021a14918b
commit
31677b1ecc
4 changed files with 145 additions and 4 deletions
|
|
@ -385,9 +385,23 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
return
|
||||
|
||||
await self._log_batch_to_rubrik(
|
||||
data=self.log_queue,
|
||||
data=list(self.log_queue),
|
||||
)
|
||||
|
||||
async def flush_queue(self):
|
||||
if self.flush_lock is None:
|
||||
return
|
||||
|
||||
async with self.flush_lock:
|
||||
if self.log_queue:
|
||||
log_queue_snapshot = list(self.log_queue)
|
||||
verbose_logger.debug(
|
||||
"Rubrik: Flushing batch of %s events", len(log_queue_snapshot)
|
||||
)
|
||||
await self._log_batch_to_rubrik(data=log_queue_snapshot)
|
||||
del self.log_queue[: len(log_queue_snapshot)]
|
||||
self.last_flush_time = time.time()
|
||||
|
||||
# -- Tool blocking service -------------------------------------------------
|
||||
|
||||
async def _post_to_tool_blocking_service(
|
||||
|
|
|
|||
|
|
@ -687,6 +687,62 @@ class VertexBase:
|
|||
# Re-raise the original error for better context
|
||||
raise error
|
||||
|
||||
async def _handle_reauthentication_async(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
project_id: Optional[str],
|
||||
credential_cache_key: Tuple,
|
||||
error: Exception,
|
||||
) -> Tuple[str, str]:
|
||||
"""
|
||||
Async reauthentication retry that stays within the per-key async lock.
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
f"Handling async reauthentication for project_id: {project_id}. "
|
||||
f"Clearing cache and retrying once."
|
||||
)
|
||||
|
||||
self._credentials_project_mapping.pop(credential_cache_key, None)
|
||||
|
||||
try:
|
||||
_credentials, credential_project_id = (
|
||||
await self._load_and_cache_credentials(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
)
|
||||
)
|
||||
if project_id is None and isinstance(credential_project_id, str):
|
||||
project_id = credential_project_id
|
||||
cache_credentials = (
|
||||
json.dumps(credentials)
|
||||
if isinstance(credentials, dict)
|
||||
else credentials
|
||||
)
|
||||
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 _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
|
||||
except Exception as retry_error:
|
||||
verbose_logger.error(
|
||||
f"Async reauthentication retry failed for project_id: {project_id}. "
|
||||
f"Original error: {str(error)}. Retry error: {str(retry_error)}"
|
||||
)
|
||||
raise error
|
||||
|
||||
def get_access_token(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
|
|
@ -935,9 +991,7 @@ class VertexBase:
|
|||
verbose_logger.debug(
|
||||
"Reauthentication needed, clearing cache and retrying"
|
||||
)
|
||||
if credential_cache_key in self._credentials_project_mapping:
|
||||
del self._credentials_project_mapping[credential_cache_key]
|
||||
return await asyncify(self._handle_reauthentication)(
|
||||
return await self._handle_reauthentication_async(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
|
|
|
|||
|
|
@ -220,6 +220,22 @@ class TestBatchLogging:
|
|||
handler.async_httpx_client.post.assert_called_once()
|
||||
assert len(handler.log_queue) == 0
|
||||
|
||||
async def test_flush_queue_preserves_events_added_during_send(self, handler):
|
||||
handler.log_queue = [{"msg": "a"}, {"msg": "b"}]
|
||||
|
||||
async def mock_post(*_args, **_kwargs):
|
||||
handler.log_queue.append({"msg": "c"})
|
||||
mock_response = Mock()
|
||||
mock_response.raise_for_status = Mock()
|
||||
return mock_response
|
||||
|
||||
handler.async_httpx_client = AsyncMock()
|
||||
handler.async_httpx_client.post = mock_post
|
||||
|
||||
await handler.flush_queue()
|
||||
|
||||
assert handler.log_queue == [{"msg": "c"}]
|
||||
|
||||
async def test_log_batch_error_does_not_crash(self, handler):
|
||||
handler.log_queue = [{"msg": "a"}]
|
||||
mock_response = Mock()
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
|
@ -1512,6 +1513,62 @@ class TestVertexBase:
|
|||
refresh_call_count == 1
|
||||
), f"Expected 1 refresh call, got {refresh_call_count}"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_reauthentication_uses_async_single_flight(self):
|
||||
"""Concurrent async reauth should reload once without using the sync path."""
|
||||
from google.auth.credentials import TokenState
|
||||
|
||||
vertex_base = VertexBase()
|
||||
stale_creds = MagicMock()
|
||||
stale_creds.token = "expired-token"
|
||||
stale_creds.token_state = TokenState.INVALID
|
||||
stale_creds.project_id = "project-1"
|
||||
stale_creds.quota_project_id = "project-1"
|
||||
|
||||
refreshed_creds = MagicMock()
|
||||
refreshed_creds.token = "refreshed-token"
|
||||
refreshed_creds.token_state = TokenState.FRESH
|
||||
refreshed_creds.project_id = "project-1"
|
||||
refreshed_creds.quota_project_id = "project-1"
|
||||
|
||||
credentials = {"type": "service_account", "project_id": "project-1"}
|
||||
cache_key = (json.dumps(credentials), "project-1")
|
||||
vertex_base._credentials_project_mapping[cache_key] = (
|
||||
stale_creds,
|
||||
"project-1",
|
||||
)
|
||||
|
||||
load_call_count = 0
|
||||
|
||||
def load_auth_impl(*_args, **_kwargs):
|
||||
nonlocal load_call_count
|
||||
load_call_count += 1
|
||||
return refreshed_creds, "project-1"
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
vertex_base,
|
||||
"refresh_auth",
|
||||
side_effect=Exception("Reauthentication is needed"),
|
||||
),
|
||||
patch.object(vertex_base, "load_auth", side_effect=load_auth_impl),
|
||||
patch.object(vertex_base, "get_access_token") as mock_get_access_token,
|
||||
):
|
||||
results = await asyncio.gather(
|
||||
*[
|
||||
vertex_base._ensure_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id="project-1",
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
for _ in range(10)
|
||||
]
|
||||
)
|
||||
|
||||
assert results == [("refreshed-token", "project-1")] * 10
|
||||
assert load_call_count == 1
|
||||
mock_get_access_token.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_refresh_when_near_expiry(self):
|
||||
"""When token_state is STALE (within the 3:45 REFRESH_THRESHOLD window),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue