Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
118 changes: 54 additions & 64 deletions agentplatform/_genai/_evals_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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"}


Expand Down Expand Up @@ -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"}


Expand Down Expand Up @@ -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"}


Expand Down
14 changes: 14 additions & 0 deletions agentplatform/_genai/_evals_constant.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
)
17 changes: 2 additions & 15 deletions agentplatform/_genai/_evals_metric_handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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.
Expand All @@ -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"
Expand Down
132 changes: 114 additions & 18 deletions tests/unit/agentplatform/genai/test_evals.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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 = [
{
Expand All @@ -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",
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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."""

Expand Down
Loading