fix: Add AuthConfig json serialization in vertex ai session service

Resolves https://github.com/google/adk-python/issues/2035

PiperOrigin-RevId: 816387628
This commit is contained in:
Shangjie Chen
2025-10-07 14:47:28 -07:00
committed by Copybara-Service
parent 3f86760e0b
commit 636def3687
2 changed files with 33 additions and 67 deletions
@@ -14,11 +14,11 @@
from __future__ import annotations from __future__ import annotations
import datetime import datetime
import json
import logging import logging
import os import os
import re import re
from typing import Any from typing import Any
from typing import Dict
from typing import Optional from typing import Optional
from typing import Union from typing import Union
@@ -132,9 +132,9 @@ class VertexAiSessionService(BaseSessionService):
session_id = get_session_response.name.split('/')[-1] session_id = get_session_response.name.split('/')[-1]
session = Session( session = Session(
app_name=str(app_name), app_name=app_name,
user_id=str(user_id), user_id=user_id,
id=str(session_id), id=session_id,
state=getattr(get_session_response, 'session_state', None) or {}, state=getattr(get_session_response, 'session_state', None) or {},
last_update_time=get_session_response.update_time.timestamp(), last_update_time=get_session_response.update_time.timestamp(),
) )
@@ -162,12 +162,11 @@ class VertexAiSessionService(BaseSessionService):
f'Session {session_id} does not belong to user {user_id}.' f'Session {session_id} does not belong to user {user_id}.'
) )
session_id_from_name = get_session_response.name.split('/')[-1]
update_timestamp = get_session_response.update_time.timestamp() update_timestamp = get_session_response.update_time.timestamp()
session = Session( session = Session(
app_name=str(app_name), app_name=app_name,
user_id=str(user_id), user_id=user_id,
id=str(session_id_from_name), id=session_id,
state=getattr(get_session_response, 'session_state', None) or {}, state=getattr(get_session_response, 'session_state', None) or {},
last_update_time=update_timestamp, last_update_time=update_timestamp,
) )
@@ -186,10 +185,10 @@ class VertexAiSessionService(BaseSessionService):
name=f'reasoningEngines/{reasoning_engine_id}/sessions/{session_id}', name=f'reasoningEngines/{reasoning_engine_id}/sessions/{session_id}',
**list_events_kwargs, **list_events_kwargs,
) )
session.events += [_from_api_event(event) for event in events_iterator] session.events += [
_from_api_event(event)
session.events = [ for event in events_iterator
event for event in session.events if event.timestamp <= update_timestamp if event.timestamp.timestamp() <= update_timestamp
] ]
# Filter events based on config # Filter events based on config
@@ -259,7 +258,12 @@ class VertexAiSessionService(BaseSessionService):
'artifact_delta': event.actions.artifact_delta, 'artifact_delta': event.actions.artifact_delta,
'transfer_agent': event.actions.transfer_to_agent, 'transfer_agent': event.actions.transfer_to_agent,
'escalate': event.actions.escalate, 'escalate': event.actions.escalate,
'requested_auth_configs': event.actions.requested_auth_configs, 'requested_auth_configs': {
k: json.loads(v.model_dump_json(exclude_none=True, by_alias=True))
for k, v in event.actions.requested_auth_configs.items()
},
# TODO: add requested_tool_confirmations, compaction, agent_state once
# they are available in the API.
} }
if event.error_code: if event.error_code:
config['error_code'] = event.error_code config['error_code'] = event.error_code
@@ -305,7 +309,7 @@ class VertexAiSessionService(BaseSessionService):
pattern = r'^projects/([a-zA-Z0-9-_]+)/locations/([a-zA-Z0-9-_]+)/reasoningEngines/(\d+)$' pattern = r'^projects/([a-zA-Z0-9-_]+)/locations/([a-zA-Z0-9-_]+)/reasoningEngines/(\d+)$'
match = re.fullmatch(pattern, app_name) match = re.fullmatch(pattern, app_name)
if not bool(match): if not match:
raise ValueError( raise ValueError(
f'App name {app_name} is not valid. It should either be the full' f'App name {app_name} is not valid. It should either be the full'
' ReasoningEngine resource name, or the reasoning engine id.' ' ReasoningEngine resource name, or the reasoning engine id.'
@@ -342,59 +346,6 @@ def _is_vertex_express_mode(
) )
def _convert_event_to_json(event: Event) -> Dict[str, Any]:
metadata_json = {
'partial': event.partial,
'turn_complete': event.turn_complete,
'interrupted': event.interrupted,
'branch': event.branch,
'custom_metadata': event.custom_metadata,
'long_running_tool_ids': (
list(event.long_running_tool_ids)
if event.long_running_tool_ids
else None
),
}
if event.grounding_metadata:
metadata_json['grounding_metadata'] = event.grounding_metadata.model_dump(
exclude_none=True, mode='json'
)
event_json = {
'author': event.author,
'invocation_id': event.invocation_id,
'timestamp': {
'seconds': int(event.timestamp),
'nanos': int(
(event.timestamp - int(event.timestamp)) * 1_000_000_000
),
},
'error_code': event.error_code,
'error_message': event.error_message,
'event_metadata': metadata_json,
}
if event.actions:
actions_json = {
'skip_summarization': event.actions.skip_summarization,
'state_delta': event.actions.state_delta,
'artifact_delta': event.actions.artifact_delta,
'transfer_agent': event.actions.transfer_to_agent,
'escalate': event.actions.escalate,
'requested_auth_configs': event.actions.requested_auth_configs,
}
event_json['actions'] = actions_json
if event.content:
event_json['content'] = event.content.model_dump(
exclude_none=True, mode='json'
)
if event.error_code:
event_json['error_code'] = event.error_code
if event.error_message:
event_json['error_message'] = event.error_message
return event_json
def _from_api_event(api_event_obj: vertexai.types.SessionEvent) -> Event: def _from_api_event(api_event_obj: vertexai.types.SessionEvent) -> Event:
"""Converts an API event object to an Event object.""" """Converts an API event object to an Event object."""
actions = getattr(api_event_obj, 'actions', None) actions = getattr(api_event_obj, 'actions', None)
@@ -22,6 +22,9 @@ from typing import Tuple
from unittest import mock from unittest import mock
from dateutil.parser import isoparse from dateutil.parser import isoparse
from fastapi.openapi import models as openapi_models
from google.adk.auth import auth_schemes
from google.adk.auth.auth_tool import AuthConfig
from google.adk.events.event import Event from google.adk.events.event import Event
from google.adk.events.event_actions import EventActions from google.adk.events.event_actions import EventActions
from google.adk.sessions.base_session_service import GetSessionConfig from google.adk.sessions.base_session_service import GetSessionConfig
@@ -522,6 +525,18 @@ async def test_append_event():
transfer_to_agent='another_agent', transfer_to_agent='another_agent',
state_delta={'new_key': 'new_value'}, state_delta={'new_key': 'new_value'},
skip_summarization=True, skip_summarization=True,
requested_auth_configs={
'test_auth': AuthConfig(
auth_scheme=auth_schemes.OAuth2(
flows=openapi_models.OAuthFlows(
implicit=openapi_models.OAuthFlowImplicit(
authorizationUrl='http://test.com/auth',
scopes={},
)
)
),
),
},
), ),
error_code='1', error_code='1',
error_message='test_error', error_message='test_error',