fix: preserve async batch and reauth coordination

This commit is contained in:
Cursor Agent 2026-05-04 13:44:59 +00:00
parent 021a14918b
commit 31677b1ecc
No known key found for this signature in database
4 changed files with 145 additions and 4 deletions

View file

@ -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(

View file

@ -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,

View file

@ -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()

View file

@ -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),