From f6bf2195cd1503b1a41586cc617b6939941f5147 Mon Sep 17 00:00:00 2001 From: Jason Dai Date: Wed, 7 Oct 2026 15:33:31 -0700 Subject: [PATCH] fix: GenAI Client(evals) - Back off only on retryable errors in inference and agent-run retries PiperOrigin-RevId: 995397403 --- agentplatform/_genai/_evals_common.py | 118 +++++++--------- agentplatform/_genai/_evals_constant.py | 14 ++ .../_genai/_evals_metric_handlers.py | 17 +-- tests/unit/agentplatform/genai/test_evals.py | 132 +++++++++++++++--- 4 files changed, 184 insertions(+), 97 deletions(-) diff --git a/agentplatform/_genai/_evals_common.py b/agentplatform/_genai/_evals_common.py index d21bed6d07..85c3ca6bd9 100644 --- a/agentplatform/_genai/_evals_common.py +++ b/agentplatform/_genai/_evals_common.py @@ -22,6 +22,7 @@ import json import logging import os +import random import threading import time from typing import Any, Callable, Literal, Optional, Union, cast @@ -160,6 +161,45 @@ def _get_runtime_instance( return _thread_local_data.runtime_instances[agent_name] +def _is_retryable_error(e: Exception) -> bool: + """Returns True for transient errors that are safe to retry with backoff.""" + if isinstance(e, genai_errors.APIError): + return e.code in _evals_constant.RETRYABLE_STATUS_CODES + return isinstance( + e, (api_exceptions.ResourceExhausted, api_exceptions.ServiceUnavailable) + ) + + +def _get_retry_backoff( + e: Exception, attempt: int, max_retries: int, operation: str +) -> Optional[float]: + """Returns the seconds to wait before the next attempt, or None to stop. + + Retryable errors back off exponentially with jitter. Other errors fail + fast, and the last attempt never waits. + """ + if not _is_retryable_error(e) or attempt == max_retries - 1: + logger.error( + "Error during %s on attempt %d/%d: %s", + operation, + attempt + 1, + max_retries, + e, + ) + return None + backoff = 2**attempt + random.uniform(0, 1) + logger.warning( + "Retryable error during %s on attempt %d/%d: %s. Retrying in %.1f" + " seconds...", + operation, + attempt + 1, + max_retries, + e, + backoff, + ) + return backoff + + def _generate_content_with_retry( api_client: BaseApiClient, model: str, @@ -222,29 +262,11 @@ def _generate_content_with_retry( } else: return response - except api_exceptions.ResourceExhausted as e: - logger.warning( - "Resource Exhausted error on attempt %d/%d: %s. Retrying in %s" - " seconds...", - attempt + 1, - max_retries, - e, - 2**attempt, - ) - if attempt == max_retries - 1: - return {"error": f"Resource exhausted after retries: {e}"} - time.sleep(2**attempt) except Exception as e: # pylint: disable=broad-exception-caught - logger.error( - "Unexpected error during generate_content on attempt %d/%d: %s", - attempt + 1, - max_retries, - e, - ) - - if attempt == max_retries - 1: - return {"error": f"Failed after retries: {e}"} - time.sleep(1) + backoff = _get_retry_backoff(e, attempt, max_retries, "generate_content") + if backoff is None: + return {"error": f"Failed on attempt {attempt + 1}/{max_retries}: {e}"} + time.sleep(backoff) return {"error": f"Failed to generate content after {max_retries} retries"} @@ -3570,28 +3592,11 @@ def _execute_agent_run_with_retry( if event and CONTENT in event and PARTS in event[CONTENT]: responses.append(event) return responses - except api_exceptions.ResourceExhausted as e: - logger.warning( - "Resource Exhausted error on attempt %d/%d: %s. Retrying in %s" - " seconds...", - attempt + 1, - max_retries, - e, - 2**attempt, - ) - if attempt == max_retries - 1: - return {"error": f"Resource exhausted after retries: {e}"} - time.sleep(2**attempt) except Exception as e: # pylint: disable=broad-exception-caught - logger.error( - "Unexpected error during agent engine run on attempt %d/%d: %s", - attempt + 1, - max_retries, - e, - ) - if attempt == max_retries - 1: - return {"error": f"Failed after retries: {e}"} - time.sleep(1) + backoff = _get_retry_backoff(e, attempt, max_retries, "agent engine run") + if backoff is None: + return {"error": f"Failed on attempt {attempt + 1}/{max_retries}: {e}"} + time.sleep(backoff) return {"error": f"Failed to get agent run results after {max_retries} retries"} @@ -3697,28 +3702,13 @@ async def _execute_local_agent_run_with_retry_async( if event and CONTENT in event and PARTS in event[CONTENT]: events.append(event) return events - except api_exceptions.ResourceExhausted as e: - logger.warning( - "Resource Exhausted error on attempt %d/%d: %s. Retrying" - " in %s seconds...", - attempt + 1, - max_retries, - e, - 2**attempt, - ) - if attempt == max_retries - 1: - return {"error": f"Resource exhausted after retries: {e}"} - await asyncio.sleep(2**attempt) except Exception as e: # pylint: disable=broad-exception-caught - logger.error( - "Unexpected error during agent run on attempt %d/%d: %s", - attempt + 1, - max_retries, - e, - ) - if attempt == max_retries - 1: - return {"error": f"Failed after retries: {e}"} - await asyncio.sleep(1) + backoff = _get_retry_backoff(e, attempt, max_retries, "agent run") + if backoff is None: + return { + "error": f"Failed on attempt {attempt + 1}/{max_retries}: {e}" + } + await asyncio.sleep(backoff) return {"error": f"Failed to get agent run results after {max_retries} retries"} diff --git a/agentplatform/_genai/_evals_constant.py b/agentplatform/_genai/_evals_constant.py index 8609c05748..0decd21793 100644 --- a/agentplatform/_genai/_evals_constant.py +++ b/agentplatform/_genai/_evals_constant.py @@ -120,3 +120,17 @@ INTERACTION_ID, } ) + +# HTTP status codes that are safe to retry with backoff. +RETRYABLE_STATUS_CODES = frozenset( + { + 408, # RequestTimeout (DEADLINE_EXCEEDED) + 409, # Conflict / Aborted (ABORTED) + 429, # TooManyRequests / ResourceExhausted (RESOURCE_EXHAUSTED) + 499, # Client Closed Request (CANCELLED) + 500, # InternalServerError (INTERNAL) + 502, # BadGateway + 503, # ServiceUnavailable (UNAVAILABLE) + 504, # GatewayTimeout (DEADLINE_EXCEEDED) + } +) diff --git a/agentplatform/_genai/_evals_metric_handlers.py b/agentplatform/_genai/_evals_metric_handlers.py index 77f5db46e4..2646735de3 100644 --- a/agentplatform/_genai/_evals_metric_handlers.py +++ b/agentplatform/_genai/_evals_metric_handlers.py @@ -39,19 +39,6 @@ logger = logging.getLogger(__name__) _MAX_RETRIES = 5 -# HTTP status codes that are safe to retry with backoff. -_RETRYABLE_STATUS_CODES = frozenset( - { - 408, # RequestTimeout (DEADLINE_EXCEEDED) - 409, # Conflict / Aborted (ABORTED) - 429, # TooManyRequests / ResourceExhausted (RESOURCE_EXHAUSTED) - 499, # Client Closed Request (CANCELLED) - 500, # InternalServerError (INTERNAL) - 502, # BadGateway - 503, # ServiceUnavailable (UNAVAILABLE) - 504, # GatewayTimeout (DEADLINE_EXCEEDED) - } -) R = TypeVar("R") T = TypeVar("T", types.Metric, types.MetricSource, types.LLMMetric) @@ -64,7 +51,7 @@ def _call_with_retry( """Calls ``fn()`` with exponential backoff + jitter on retryable errors. Retries up to ``_MAX_RETRIES`` times on errors whose HTTP status code is - in ``_RETRYABLE_STATUS_CODES`` (Aborted, DeadlineExceeded, + in ``_evals_constant.RETRYABLE_STATUS_CODES`` (Aborted, DeadlineExceeded, ResourceExhausted, ServiceUnavailable, Cancelled). Non-retryable errors are re-raised immediately. If all retries are exhausted the last exception is re-raised so the caller can decide how to handle it. @@ -84,7 +71,7 @@ def _call_with_retry( try: return fn() except genai_errors.APIError as e: - if e.code in _RETRYABLE_STATUS_CODES: + if e.code in _evals_constant.RETRYABLE_STATUS_CODES: backoff = 2**attempt + random.uniform(0, 1) logger.warning( "Retryable error (code=%s) on attempt %d/%d for metric" diff --git a/tests/unit/agentplatform/genai/test_evals.py b/tests/unit/agentplatform/genai/test_evals.py index a6976942ba..492d9cb429 100644 --- a/tests/unit/agentplatform/genai/test_evals.py +++ b/tests/unit/agentplatform/genai/test_evals.py @@ -26,6 +26,7 @@ import tempfile from unittest import mock +from google.api_core import exceptions as api_exceptions import google.auth.credentials from google.cloud import aiplatform import agentplatform @@ -3390,9 +3391,7 @@ def test_run_inference_with_runtime_and_session_inputs_dict( ] mock_runtime.stream_query.return_value = iter(stream_query_return_value) - mock_agentplatform_client.return_value.runtimes.get.return_value = ( - mock_runtime - ) + mock_agentplatform_client.return_value.runtimes.get.return_value = mock_runtime inference_result = self.client.evals.run_inference( agent="projects/test-project/locations/us-central1/reasoningEngines/123", @@ -3499,9 +3498,7 @@ def test_run_inference_with_runtime_and_session_inputs_literal_string( ] mock_runtime.stream_query.return_value = iter(stream_query_return_value) - mock_agentplatform_client.return_value.runtimes.get.return_value = ( - mock_runtime - ) + mock_agentplatform_client.return_value.runtimes.get.return_value = mock_runtime inference_result = self.client.evals.run_inference( agent="projects/test-project/locations/us-central1/reasoningEngines/123", @@ -3591,9 +3588,7 @@ def test_run_inference_with_runtime_with_response_column_raises_error( ) mock_runtime = mock.Mock() - mock_agentplatform_client.return_value.runtimes.get.return_value = ( - mock_runtime - ) + mock_agentplatform_client.return_value.runtimes.get.return_value = mock_runtime with pytest.raises(ValueError) as excinfo: self.client.evals.run_inference( @@ -3644,9 +3639,7 @@ def test_run_inference_with_runtime_falls_back_to_managed_sessions_api( "projects/test-project/locations/us-central1" "/reasoningEngines/123/sessions/managed-session-1" ) - mock_runtime.api_client.sessions.create.return_value = ( - mock_session_operation - ) + mock_runtime.api_client.sessions.create.return_value = mock_session_operation stream_query_return_value = [ { @@ -3663,9 +3656,7 @@ def test_run_inference_with_runtime_falls_back_to_managed_sessions_api( }, ] mock_runtime.stream_query.return_value = iter(stream_query_return_value) - mock_agentplatform_client.return_value.runtimes.get.return_value = ( - mock_runtime - ) + mock_agentplatform_client.return_value.runtimes.get.return_value = mock_runtime inference_result = self.client.evals.run_inference( agent="projects/test-project/locations/us-central1/reasoningEngines/123", @@ -4392,9 +4383,7 @@ def test_run_inference_non_gemini_agent_routes_to_runtime( } ] ) - mock_agentplatform_client.return_value.runtimes.get.return_value = ( - mock_runtime - ) + mock_agentplatform_client.return_value.runtimes.get.return_value = mock_runtime self.client.evals.run_inference( src=mock_df, @@ -11141,6 +11130,113 @@ def test_call_with_retry_no_retry_on_non_retryable(self, mock_sleep): assert mock_sleep.call_count == 0 +class TestInferenceRetry: + @pytest.mark.parametrize( + "error, expected", + [ + (genai_errors.ClientError(code=429, response_json={}), True), + (genai_errors.ServerError(code=503, response_json={}), True), + (genai_errors.ClientError(code=400, response_json={}), False), + (genai_errors.ClientError(code=403, response_json={}), False), + (genai_errors.ClientError(code=404, response_json={}), False), + (api_exceptions.ResourceExhausted("quota"), True), + (api_exceptions.ServiceUnavailable("unavailable"), True), + (api_exceptions.InvalidArgument("bad request"), False), + (ValueError("bad value"), False), + ], + ) + def test_is_retryable_error(self, error, expected): + assert _evals_common._is_retryable_error(error) is expected + + @mock.patch("time.sleep", return_value=None) + @mock.patch.object(_evals_common, "Models") + def test_generate_content_retries_429_with_backoff(self, mock_models, mock_sleep): + response = genai_types.GenerateContentResponse( + candidates=[ + genai_types.Candidate( + content=genai_types.Content(parts=[genai_types.Part(text="ok")]), + finish_reason=genai_types.FinishReason.STOP, + ) + ] + ) + error = genai_errors.ClientError(code=429, response_json={}) + mock_models.return_value.generate_content.side_effect = [error, error, response] + + result = _evals_common._generate_content_with_retry( + api_client=mock.Mock(), model="gemini-pro", contents="prompt" + ) + + assert result is response + assert mock_models.return_value.generate_content.call_count == 3 + backoffs = [call.args[0] for call in mock_sleep.call_args_list] + assert len(backoffs) == 2 + assert 1 <= backoffs[0] <= 2 + assert 2 <= backoffs[1] <= 3 + + @mock.patch("time.sleep", return_value=None) + @mock.patch.object(_evals_common, "Models") + def test_generate_content_does_not_retry_400(self, mock_models, mock_sleep): + mock_models.return_value.generate_content.side_effect = ( + genai_errors.ClientError(code=400, response_json={}) + ) + + result = _evals_common._generate_content_with_retry( + api_client=mock.Mock(), model="gemini-pro", contents="prompt" + ) + + assert result["error"].startswith("Failed on attempt 1/3: 400") + assert mock_models.return_value.generate_content.call_count == 1 + mock_sleep.assert_not_called() + + @mock.patch("time.sleep", return_value=None) + def test_agent_engine_run_retries_429_and_fails_fast_on_400(self, mock_sleep): + runtime = mock.Mock() + runtime.stream_query.side_effect = [ + genai_errors.ClientError(code=429, response_json={}), + genai_errors.ClientError(code=400, response_json={}), + ] + + with mock.patch.object( + _evals_common, "_create_runtime_session", return_value="session-id" + ): + result = _evals_common._execute_agent_run_with_retry( + row=pd.Series({"prompt": "prompt"}), contents="prompt", runtime=runtime + ) + + assert result["error"].startswith("Failed on attempt 2/3: 400") + assert runtime.stream_query.call_count == 2 + mock_sleep.assert_called_once() + assert 1 <= mock_sleep.call_args.args[0] <= 2 + + @mock.patch("asyncio.sleep", new_callable=mock.AsyncMock) + def test_local_agent_run_retries_429_and_fails_fast_on_400(self, mock_sleep): + mock_runners = mock.MagicMock() + mock_runners.Runner.return_value.run_async.side_effect = [ + genai_errors.ClientError(code=429, response_json={}), + genai_errors.ClientError(code=400, response_json={}), + ] + mock_sessions = mock.MagicMock() + mock_sessions.InMemorySessionService.return_value.create_session = ( + mock.AsyncMock() + ) + + with mock.patch.dict( + sys.modules, + {"google.adk.runners": mock_runners, "google.adk.sessions": mock_sessions}, + ): + result = _evals_common._execute_local_agent_run_with_retry( + row=pd.Series({"prompt": "prompt"}), + contents="prompt", + agent=mock.Mock(), + api_client=mock.Mock(), + ) + + assert result["error"].startswith("Failed on attempt 2/3: 400") + assert mock_runners.Runner.return_value.run_async.call_count == 2 + mock_sleep.assert_awaited_once() + assert 1 <= mock_sleep.await_args.args[0] <= 2 + + class TestComputationMetricRetry: """Tests for retry behavior in ComputationMetricHandler."""