mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
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:
committed by
Copybara-Service
parent
3f86760e0b
commit
636def3687
@@ -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',
|
||||||
|
|||||||
Reference in New Issue
Block a user