feat: Remove custom polling logic for Vertex AI Session Service since LRO polling is supported in express mode

PiperOrigin-RevId: 826226731
This commit is contained in:
Google Team Member
2025-10-30 16:13:43 -07:00
committed by Copybara-Service
parent dea7668d1a
commit 546c2a6816
3 changed files with 8 additions and 86 deletions
@@ -35,7 +35,6 @@ from . import _session_util
from ..events.event import Event
from ..events.event_actions import EventActions
from ..utils.vertex_ai_utils import get_express_mode_api_key
from ..utils.vertex_ai_utils import is_vertex_express_mode
from .base_session_service import BaseSessionService
from .base_session_service import GetSessionConfig
from .base_session_service import ListSessionsResponse
@@ -115,53 +114,14 @@ class VertexAiSessionService(BaseSessionService):
config = {'session_state': state} if state else {}
config.update(kwargs)
if is_vertex_express_mode(
self._project, self._location, self._express_mode_api_key
):
config['wait_for_completion'] = False
api_response = await api_client.aio.agent_engines.sessions.create(
name=f'reasoningEngines/{reasoning_engine_id}',
user_id=user_id,
config=config,
)
logger.info('Create session response received.')
session_id = api_response.name.split('/')[-3]
# Express mode doesn't support LRO, so we need to poll
# the session resource.
# TODO: remove this once LRO polling is supported in Express mode.
@retry(
stop=stop_after_attempt(6),
wait=wait_exponential(multiplier=1, min=1, max=3),
retry=retry_if_result(lambda response: not response),
reraise=True,
)
async def _poll_session_resource():
try:
return await api_client.aio.agent_engines.sessions.get(
name=f'reasoningEngines/{reasoning_engine_id}/sessions/{session_id}'
)
except ClientError:
logger.info('Polling session resource')
return None
try:
await _poll_session_resource()
except Exception as exc:
raise ValueError('Failed to create session.') from exc
get_session_response = await api_client.aio.agent_engines.sessions.get(
name=f'reasoningEngines/{reasoning_engine_id}/sessions/{session_id}'
)
else:
api_response = await api_client.aio.agent_engines.sessions.create(
name=f'reasoningEngines/{reasoning_engine_id}',
user_id=user_id,
config=config,
)
logger.debug('Create session response: %s', api_response)
get_session_response = api_response.response
session_id = get_session_response.name.split('/')[-1]
api_response = await api_client.aio.agent_engines.sessions.create(
name=f'reasoningEngines/{reasoning_engine_id}',
user_id=user_id,
config=config,
)
logger.debug('Create session response: %s', api_response)
get_session_response = api_response.response
session_id = get_session_response.name.split('/')[-1]
session = Session(
app_name=app_name,
-12
View File
@@ -26,18 +26,6 @@ from typing import Optional
from ..utils.env_utils import is_env_enabled
def is_vertex_express_mode(
project: Optional[str], location: Optional[str], api_key: Optional[str]
) -> bool:
"""Check if Vertex AI and API key are both enabled replacing project and location, meaning the user is using the Vertex Express Mode."""
return (
is_env_enabled('GOOGLE_GENAI_USE_VERTEXAI')
and api_key is not None
and project is None
and location is None
)
def get_express_mode_api_key(
project: Optional[str],
location: Optional[str],
@@ -20,32 +20,6 @@ from google.adk.utils import vertex_ai_utils
import pytest
@pytest.mark.parametrize(
('use_vertexai_env', 'project', 'location', 'api_key', 'expected'),
[
('true', None, None, 'test-key', True),
('1', None, None, 'test-key', True),
('false', None, None, 'test-key', False),
('0', None, None, 'test-key', False),
(None, None, None, 'test-key', False),
('true', 'test-project', None, 'test-key', False),
('true', None, 'test-location', 'test-key', False),
('true', None, None, None, False),
],
)
def test_is_vertex_express_mode(
use_vertexai_env, project, location, api_key, expected
):
env_vars = {}
if use_vertexai_env:
env_vars['GOOGLE_GENAI_USE_VERTEXAI'] = use_vertexai_env
with mock.patch.dict('os.environ', env_vars, clear=True):
assert (
vertex_ai_utils.is_vertex_express_mode(project, location, api_key)
== expected
)
def test_get_express_mode_api_key_value_error():
with pytest.raises(ValueError) as excinfo:
vertex_ai_utils.get_express_mode_api_key(