diff --git a/litellm/integrations/rubrik.py b/litellm/integrations/rubrik.py index 48cf165ab3a..80979e48ae8 100644 --- a/litellm/integrations/rubrik.py +++ b/litellm/integrations/rubrik.py @@ -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( diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 811519dab88..d5516a38785 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -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, diff --git a/tests/test_litellm/integrations/test_rubrik.py b/tests/test_litellm/integrations/test_rubrik.py index 3a526fed270..6e51106f632 100644 --- a/tests/test_litellm/integrations/test_rubrik.py +++ b/tests/test_litellm/integrations/test_rubrik.py @@ -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() diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py index 5a30b956b23..4aefacba643 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py @@ -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),