fix: prevent duplicate OAuth prompts and fix tool resumption
Merge https://github.com/google/adk-python/pull/5985 PiperOrigin-RevId: 967017952
This commit is contained in:
committed by
Copybara-Service
parent
0fd681e7d2
commit
eaad2f83b9
@@ -27,10 +27,53 @@ from ..flows.llm_flows.functions import handle_function_calls_async
|
||||
from ..flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME
|
||||
from ..models.llm_request import LlmRequest
|
||||
from ..sessions.state import State
|
||||
from .auth_credential import AuthCredential
|
||||
from .auth_handler import AuthHandler
|
||||
from .auth_tool import AuthConfig
|
||||
from .auth_tool import AuthToolArguments
|
||||
|
||||
|
||||
def _merge_credential_oauth2_fields(
|
||||
target_cred: AuthCredential | None,
|
||||
source_cred: AuthCredential | None,
|
||||
) -> AuthCredential | None:
|
||||
"""Merges OAuth2 fields from source_cred into target_cred if target_cred fields are None.
|
||||
|
||||
If target_cred is None, returns source_cred.
|
||||
Otherwise, merges fields and returns target_cred.
|
||||
"""
|
||||
if not source_cred:
|
||||
return target_cred
|
||||
if not target_cred:
|
||||
return source_cred
|
||||
|
||||
if target_cred.oauth2 is None and source_cred.oauth2 is not None:
|
||||
target_cred.oauth2 = source_cred.oauth2.model_copy(deep=True)
|
||||
elif target_cred.oauth2 and source_cred.oauth2:
|
||||
target = target_cred.oauth2
|
||||
source = source_cred.oauth2
|
||||
for field in [
|
||||
"client_id",
|
||||
"client_secret",
|
||||
"redirect_uri",
|
||||
"code_verifier",
|
||||
"code_challenge_method",
|
||||
]:
|
||||
if getattr(target, field) is None:
|
||||
setattr(target, field, getattr(source, field))
|
||||
|
||||
# token_endpoint_auth_method has a default value "client_secret_basic" in OAuth2Auth model.
|
||||
# We only merge it if it wasn't explicitly set in target.
|
||||
target_fields_set = getattr(target, "model_fields_set", None)
|
||||
if (
|
||||
target_fields_set is None
|
||||
or "token_endpoint_auth_method" not in target_fields_set
|
||||
):
|
||||
target.token_endpoint_auth_method = source.token_endpoint_auth_method
|
||||
|
||||
return target_cred
|
||||
|
||||
|
||||
# Prefix used by toolset auth credential IDs.
|
||||
# Auth requests with this prefix are for toolset authentication (before tool
|
||||
# listing) and don't require resuming a function call.
|
||||
@@ -80,18 +123,28 @@ async def _store_auth_and_collect_resume_targets(
|
||||
except TypeError:
|
||||
continue
|
||||
|
||||
# Step 2: Store credentials. Merge credential_key from the original
|
||||
# request into the client's auth response before storing.
|
||||
authorized_keys: set[str] = set()
|
||||
for fc_id in auth_fc_ids:
|
||||
if fc_id not in auth_responses:
|
||||
continue
|
||||
auth_config = AuthConfig.model_validate(auth_responses[fc_id])
|
||||
requested_auth_config = requested_auth_config_by_id.get(fc_id)
|
||||
if (
|
||||
requested_auth_config
|
||||
and requested_auth_config.credential_key is not None
|
||||
):
|
||||
auth_config.credential_key = requested_auth_config.credential_key
|
||||
if requested_auth_config:
|
||||
if requested_auth_config.credential_key is not None:
|
||||
auth_config.credential_key = requested_auth_config.credential_key
|
||||
if requested_auth_config.raw_auth_credential:
|
||||
auth_config.raw_auth_credential = _merge_credential_oauth2_fields(
|
||||
auth_config.raw_auth_credential,
|
||||
requested_auth_config.raw_auth_credential,
|
||||
)
|
||||
if requested_auth_config.exchanged_auth_credential:
|
||||
auth_config.exchanged_auth_credential = _merge_credential_oauth2_fields(
|
||||
auth_config.exchanged_auth_credential,
|
||||
requested_auth_config.exchanged_auth_credential,
|
||||
)
|
||||
if auth_config.credential_key:
|
||||
authorized_keys.add(auth_config.credential_key)
|
||||
|
||||
await AuthHandler(auth_config=auth_config).parse_and_store_auth_response(
|
||||
state=state
|
||||
)
|
||||
@@ -121,6 +174,25 @@ async def _store_auth_and_collect_resume_targets(
|
||||
continue
|
||||
tools_to_resume.add(args.function_call_id)
|
||||
|
||||
matching_events: list[Event] = []
|
||||
for event in events:
|
||||
actions = getattr(event, "actions", None)
|
||||
if actions and actions.requested_auth_configs:
|
||||
if any(
|
||||
fc_id in actions.requested_auth_configs for fc_id in tools_to_resume
|
||||
):
|
||||
matching_events.append(event)
|
||||
|
||||
for event in matching_events:
|
||||
actions = getattr(event, "actions", None)
|
||||
if actions and actions.requested_auth_configs:
|
||||
for (
|
||||
original_fc_id,
|
||||
config,
|
||||
) in actions.requested_auth_configs.items():
|
||||
if config.credential_key in authorized_keys:
|
||||
tools_to_resume.add(original_fc_id)
|
||||
|
||||
return tools_to_resume
|
||||
|
||||
|
||||
|
||||
@@ -337,7 +337,17 @@ def build_auth_request_event(
|
||||
parts: list[types.Part] = []
|
||||
long_running_tool_ids: set[str] = set()
|
||||
|
||||
deduplicated_requests: dict[str, AuthConfig] = {}
|
||||
seen_keys = set()
|
||||
for function_call_id, auth_config in auth_requests.items():
|
||||
key = auth_config.credential_key
|
||||
if not key:
|
||||
deduplicated_requests[function_call_id] = auth_config
|
||||
elif key not in seen_keys:
|
||||
seen_keys.add(key)
|
||||
deduplicated_requests[function_call_id] = auth_config
|
||||
|
||||
for function_call_id, auth_config in deduplicated_requests.items():
|
||||
request_id = generate_client_function_call_id()
|
||||
request_euc_function_call = types.FunctionCall(
|
||||
name=REQUEST_EUC_FUNCTION_CALL_NAME,
|
||||
|
||||
@@ -20,9 +20,16 @@ from unittest.mock import AsyncMock
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import patch
|
||||
|
||||
from fastapi.openapi.models import OAuth2
|
||||
from fastapi.openapi.models import OAuthFlowAuthorizationCode
|
||||
from fastapi.openapi.models import OAuthFlows
|
||||
from google.adk.agents.invocation_context import InvocationContext
|
||||
from google.adk.auth.auth_credential import AuthCredential
|
||||
from google.adk.auth.auth_credential import AuthCredentialTypes
|
||||
from google.adk.auth.auth_credential import OAuth2Auth
|
||||
from google.adk.auth.auth_handler import AuthHandler
|
||||
from google.adk.auth.auth_preprocessor import _AuthLlmRequestProcessor
|
||||
from google.adk.auth.auth_preprocessor import _store_auth_and_collect_resume_targets
|
||||
from google.adk.auth.auth_tool import AuthConfig
|
||||
from google.adk.auth.auth_tool import AuthToolArguments
|
||||
from google.adk.events.event import Event
|
||||
@@ -82,6 +89,8 @@ class TestAuthLlmRequestProcessor:
|
||||
"""Create a mock AuthConfig."""
|
||||
config = Mock(spec=AuthConfig)
|
||||
config.credential_key = None
|
||||
config.raw_auth_credential = None
|
||||
config.exchanged_auth_credential = None
|
||||
return config
|
||||
|
||||
@pytest.fixture
|
||||
@@ -579,3 +588,430 @@ class TestAuthLlmRequestProcessor:
|
||||
result.append(event)
|
||||
|
||||
assert result == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('google.adk.auth.auth_preprocessor.AuthHandler')
|
||||
@patch('google.adk.auth.auth_tool.AuthConfig.model_validate')
|
||||
@patch('google.adk.auth.auth_preprocessor.handle_function_calls_async')
|
||||
async def test_resumes_tools_by_credential_key(
|
||||
self,
|
||||
mock_handle_function_calls,
|
||||
mock_auth_config_validate,
|
||||
mock_auth_handler_class,
|
||||
processor,
|
||||
mock_invocation_context,
|
||||
mock_llm_request,
|
||||
):
|
||||
"""Test that tools are resumed by credential key matching."""
|
||||
# Setup auth response
|
||||
auth_config = Mock(spec=AuthConfig)
|
||||
auth_config.credential_key = 'test_cred_key'
|
||||
auth_config.raw_auth_credential = None
|
||||
auth_config.exchanged_auth_credential = None
|
||||
mock_auth_config_validate.return_value = auth_config
|
||||
|
||||
auth_response = Mock()
|
||||
auth_response.name = REQUEST_EUC_FUNCTION_CALL_NAME
|
||||
auth_response.id = 'auth_fc_id'
|
||||
auth_response.response = auth_config
|
||||
|
||||
user_event = Mock(spec=Event)
|
||||
user_event.author = 'user'
|
||||
user_event.content = Mock()
|
||||
user_event.get_function_responses.return_value = [auth_response]
|
||||
user_event.get_function_calls.return_value = []
|
||||
|
||||
# Setup system event (the one that requested auth)
|
||||
system_function_call = Mock()
|
||||
system_function_call.id = 'auth_fc_id'
|
||||
system_function_call.name = REQUEST_EUC_FUNCTION_CALL_NAME
|
||||
requested_auth_config = Mock(spec=AuthConfig)
|
||||
requested_auth_config.credential_key = 'test_cred_key'
|
||||
requested_auth_config.raw_auth_credential = None
|
||||
requested_auth_config.exchanged_auth_credential = None
|
||||
|
||||
system_function_call.args = {
|
||||
'function_call_id': 'original_fc_id_1',
|
||||
'auth_config': requested_auth_config,
|
||||
}
|
||||
|
||||
system_event = Mock(spec=Event)
|
||||
system_event.content = Mock()
|
||||
system_event.get_function_calls.return_value = [system_function_call]
|
||||
|
||||
# Setup an event with actions.requested_auth_configs
|
||||
event_with_actions = Mock(spec=Event)
|
||||
event_with_actions.content = Mock()
|
||||
event_with_actions.get_function_calls.return_value = []
|
||||
|
||||
actions = Mock()
|
||||
action_config = Mock()
|
||||
action_config.credential_key = 'test_cred_key'
|
||||
actions.requested_auth_configs = {
|
||||
'original_fc_id_1': action_config,
|
||||
'original_fc_id_2': action_config,
|
||||
}
|
||||
event_with_actions.actions = actions
|
||||
|
||||
# Setup original function call events
|
||||
original_fc_1 = Mock()
|
||||
original_fc_1.id = 'original_fc_id_1'
|
||||
original_fc_2 = Mock()
|
||||
original_fc_2.id = 'original_fc_id_2'
|
||||
|
||||
original_event = Mock(spec=Event)
|
||||
original_event.content = Mock()
|
||||
original_event.get_function_calls.return_value = [
|
||||
original_fc_1,
|
||||
original_fc_2,
|
||||
]
|
||||
|
||||
# Events in order: original -> event_with_actions -> system_event -> user_event
|
||||
mock_invocation_context.session.events = [
|
||||
original_event,
|
||||
event_with_actions,
|
||||
system_event,
|
||||
user_event,
|
||||
]
|
||||
|
||||
mock_auth_handler = Mock(spec=AuthHandler)
|
||||
mock_auth_handler.parse_and_store_auth_response = AsyncMock()
|
||||
mock_auth_handler_class.return_value = mock_auth_handler
|
||||
|
||||
mock_function_response_event = Mock(spec=Event)
|
||||
mock_handle_function_calls.return_value = mock_function_response_event
|
||||
|
||||
with patch(
|
||||
'google.adk.auth.auth_tool.AuthToolArguments.model_validate'
|
||||
) as mock_auth_tool_args_validate:
|
||||
mock_args = Mock(spec=AuthToolArguments)
|
||||
mock_args.auth_config = requested_auth_config
|
||||
mock_args.function_call_id = 'original_fc_id_1'
|
||||
mock_auth_tool_args_validate.return_value = mock_args
|
||||
|
||||
result = []
|
||||
async for event in processor.run_async(
|
||||
mock_invocation_context, mock_llm_request
|
||||
):
|
||||
result.append(event)
|
||||
|
||||
mock_handle_function_calls.assert_called_once()
|
||||
call_args = mock_handle_function_calls.call_args
|
||||
assert call_args[0][1] == original_event
|
||||
assert call_args[0][3] == {'original_fc_id_1', 'original_fc_id_2'}
|
||||
assert result == [mock_function_response_event]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('google.adk.auth.auth_preprocessor.AuthHandler')
|
||||
@patch('google.adk.auth.auth_tool.AuthConfig.model_validate')
|
||||
@patch('google.adk.auth.auth_preprocessor.handle_function_calls_async')
|
||||
async def test_does_not_resume_stale_tools_from_older_events(
|
||||
self,
|
||||
mock_handle_function_calls,
|
||||
mock_auth_config_validate,
|
||||
mock_auth_handler_class,
|
||||
processor,
|
||||
mock_invocation_context,
|
||||
mock_llm_request,
|
||||
):
|
||||
"""Test that tools from older events with matching cred key are NOT resumed."""
|
||||
# Setup auth response
|
||||
auth_config = Mock(spec=AuthConfig)
|
||||
auth_config.credential_key = 'test_cred_key'
|
||||
auth_config.raw_auth_credential = None
|
||||
auth_config.exchanged_auth_credential = None
|
||||
mock_auth_config_validate.return_value = auth_config
|
||||
|
||||
auth_response = Mock()
|
||||
auth_response.name = REQUEST_EUC_FUNCTION_CALL_NAME
|
||||
auth_response.id = 'auth_fc_id'
|
||||
auth_response.response = auth_config
|
||||
|
||||
user_event = Mock(spec=Event)
|
||||
user_event.author = 'user'
|
||||
user_event.content = Mock()
|
||||
user_event.get_function_responses.return_value = [auth_response]
|
||||
user_event.get_function_calls.return_value = []
|
||||
|
||||
# Setup system event (the one that requested auth)
|
||||
system_function_call = Mock()
|
||||
system_function_call.id = 'auth_fc_id'
|
||||
system_function_call.name = REQUEST_EUC_FUNCTION_CALL_NAME
|
||||
requested_auth_config = Mock(spec=AuthConfig)
|
||||
requested_auth_config.credential_key = 'test_cred_key'
|
||||
requested_auth_config.raw_auth_credential = None
|
||||
requested_auth_config.exchanged_auth_credential = None
|
||||
|
||||
system_function_call.args = {
|
||||
'function_call_id': 'original_fc_id_1',
|
||||
'auth_config': requested_auth_config,
|
||||
}
|
||||
|
||||
system_event = Mock(spec=Event)
|
||||
system_event.content = Mock()
|
||||
system_event.get_function_calls.return_value = [system_function_call]
|
||||
|
||||
# Setup a fresh event with actions.requested_auth_configs
|
||||
fresh_event_with_actions = Mock(spec=Event)
|
||||
fresh_event_with_actions.content = Mock()
|
||||
fresh_event_with_actions.get_function_calls.return_value = []
|
||||
actions_fresh = Mock()
|
||||
action_config_fresh = Mock()
|
||||
action_config_fresh.credential_key = 'test_cred_key'
|
||||
actions_fresh.requested_auth_configs = {
|
||||
'original_fc_id_1': action_config_fresh,
|
||||
}
|
||||
fresh_event_with_actions.actions = actions_fresh
|
||||
|
||||
# Setup an OLD event with actions.requested_auth_configs that also used test_cred_key
|
||||
old_event_with_actions = Mock(spec=Event)
|
||||
old_event_with_actions.content = Mock()
|
||||
old_event_with_actions.get_function_calls.return_value = []
|
||||
actions_old = Mock()
|
||||
action_config_old = Mock()
|
||||
action_config_old.credential_key = 'test_cred_key'
|
||||
actions_old.requested_auth_configs = {'stale_fc_id': action_config_old}
|
||||
old_event_with_actions.actions = actions_old
|
||||
|
||||
# Setup original function call events
|
||||
original_fc_1 = Mock()
|
||||
original_fc_1.id = 'original_fc_id_1'
|
||||
original_fc_stale = Mock()
|
||||
original_fc_stale.id = 'stale_fc_id'
|
||||
|
||||
original_event = Mock(spec=Event)
|
||||
original_event.content = Mock()
|
||||
original_event.get_function_calls.return_value = [
|
||||
original_fc_1,
|
||||
original_fc_stale,
|
||||
]
|
||||
|
||||
# Events in order: old_event -> original -> fresh_event -> system -> user
|
||||
mock_invocation_context.session.events = [
|
||||
old_event_with_actions,
|
||||
original_event,
|
||||
fresh_event_with_actions,
|
||||
system_event,
|
||||
user_event,
|
||||
]
|
||||
|
||||
mock_auth_handler = Mock(spec=AuthHandler)
|
||||
mock_auth_handler.parse_and_store_auth_response = AsyncMock()
|
||||
mock_auth_handler_class.return_value = mock_auth_handler
|
||||
|
||||
mock_function_response_event = Mock(spec=Event)
|
||||
mock_handle_function_calls.return_value = mock_function_response_event
|
||||
|
||||
with patch(
|
||||
'google.adk.auth.auth_tool.AuthToolArguments.model_validate'
|
||||
) as mock_auth_tool_args_validate:
|
||||
mock_args = Mock(spec=AuthToolArguments)
|
||||
mock_args.auth_config = requested_auth_config
|
||||
mock_args.function_call_id = 'original_fc_id_1'
|
||||
mock_auth_tool_args_validate.return_value = mock_args
|
||||
|
||||
result = []
|
||||
async for event in processor.run_async(
|
||||
mock_invocation_context, mock_llm_request
|
||||
):
|
||||
result.append(event)
|
||||
|
||||
mock_handle_function_calls.assert_called_once()
|
||||
call_args = mock_handle_function_calls.call_args
|
||||
assert call_args[0][1] == original_event
|
||||
# Should only resume original_fc_id_1, NOT stale_fc_id
|
||||
assert call_args[0][3] == {'original_fc_id_1'}
|
||||
assert result == [mock_function_response_event]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('google.adk.auth.auth_preprocessor.AuthHandler')
|
||||
async def test_store_auth_merges_oauth2_fields(
|
||||
self,
|
||||
mock_auth_handler_class,
|
||||
):
|
||||
"""Test that OAuth2 fields are merged from requested to stored config."""
|
||||
# Setup AuthHandler mock
|
||||
mock_auth_handler = Mock(spec=AuthHandler)
|
||||
mock_auth_handler.parse_and_store_auth_response = AsyncMock()
|
||||
mock_auth_handler_class.return_value = mock_auth_handler
|
||||
|
||||
# Create requested auth config (the one in the event history)
|
||||
# It has all OAuth2 fields populated.
|
||||
requested_oauth2 = OAuth2Auth(
|
||||
client_id='expected_client_id',
|
||||
client_secret='expected_client_secret',
|
||||
redirect_uri='expected_redirect_uri',
|
||||
code_verifier='expected_code_verifier',
|
||||
code_challenge_method='S256',
|
||||
token_endpoint_auth_method='client_secret_post',
|
||||
)
|
||||
requested_auth_config = AuthConfig(
|
||||
auth_scheme=OAuth2(
|
||||
flows=OAuthFlows(
|
||||
authorizationCode=OAuthFlowAuthorizationCode(
|
||||
authorizationUrl='https://example.com/auth',
|
||||
tokenUrl='https://example.com/token',
|
||||
)
|
||||
)
|
||||
),
|
||||
raw_auth_credential=AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
oauth2=requested_oauth2,
|
||||
),
|
||||
exchanged_auth_credential=AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
oauth2=requested_oauth2,
|
||||
),
|
||||
credential_key='test_cred_key',
|
||||
)
|
||||
|
||||
# Create the auth response (the one returned by the client)
|
||||
# It has some missing OAuth2 fields that should be merged.
|
||||
stored_oauth2_raw = OAuth2Auth(
|
||||
client_id=None,
|
||||
client_secret=None,
|
||||
redirect_uri=None,
|
||||
code_verifier=None,
|
||||
code_challenge_method=None,
|
||||
access_token='some_access_token',
|
||||
)
|
||||
stored_oauth2_exchanged = OAuth2Auth(
|
||||
client_id=None,
|
||||
client_secret=None,
|
||||
redirect_uri=None,
|
||||
code_verifier=None,
|
||||
code_challenge_method=None,
|
||||
access_token='some_exchanged_token',
|
||||
)
|
||||
stored_auth_config = AuthConfig(
|
||||
auth_scheme=OAuth2(
|
||||
flows=OAuthFlows(
|
||||
authorizationCode=OAuthFlowAuthorizationCode(
|
||||
authorizationUrl='https://example.com/auth',
|
||||
tokenUrl='https://example.com/token',
|
||||
)
|
||||
)
|
||||
),
|
||||
raw_auth_credential=AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
oauth2=stored_oauth2_raw,
|
||||
),
|
||||
exchanged_auth_credential=AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
oauth2=stored_oauth2_exchanged,
|
||||
),
|
||||
credential_key='test_cred_key',
|
||||
)
|
||||
|
||||
# Setup function call in history that requested auth
|
||||
system_function_call = Mock()
|
||||
system_function_call.id = 'auth_fc_id'
|
||||
system_function_call.name = REQUEST_EUC_FUNCTION_CALL_NAME
|
||||
system_function_call.args = {
|
||||
'function_call_id': 'original_fc_id',
|
||||
'auth_config': requested_auth_config,
|
||||
}
|
||||
|
||||
system_event = Mock(spec=Event)
|
||||
system_event.content = Mock()
|
||||
system_event.get_function_calls.return_value = [system_function_call]
|
||||
|
||||
# Setup state
|
||||
mock_state = Mock()
|
||||
|
||||
# Call _store_auth_and_collect_resume_targets
|
||||
await _store_auth_and_collect_resume_targets(
|
||||
events=[system_event],
|
||||
auth_fc_ids={'auth_fc_id'},
|
||||
auth_responses={
|
||||
'auth_fc_id': stored_auth_config.model_dump(
|
||||
mode='json', exclude_defaults=True
|
||||
)
|
||||
},
|
||||
state=mock_state,
|
||||
)
|
||||
|
||||
# Verify AuthHandler was called with merged config
|
||||
mock_auth_handler_class.assert_called_once()
|
||||
called_config = mock_auth_handler_class.call_args.kwargs['auth_config']
|
||||
|
||||
# Check raw_auth_credential fields
|
||||
assert (
|
||||
called_config.raw_auth_credential.oauth2.client_id
|
||||
== 'expected_client_id'
|
||||
)
|
||||
assert (
|
||||
called_config.raw_auth_credential.oauth2.client_secret
|
||||
== 'expected_client_secret'
|
||||
)
|
||||
assert (
|
||||
called_config.raw_auth_credential.oauth2.redirect_uri
|
||||
== 'expected_redirect_uri'
|
||||
)
|
||||
assert (
|
||||
called_config.raw_auth_credential.oauth2.code_verifier
|
||||
== 'expected_code_verifier'
|
||||
)
|
||||
assert (
|
||||
called_config.raw_auth_credential.oauth2.code_challenge_method == 'S256'
|
||||
)
|
||||
assert (
|
||||
called_config.raw_auth_credential.oauth2.token_endpoint_auth_method
|
||||
== 'client_secret_post'
|
||||
)
|
||||
assert (
|
||||
called_config.raw_auth_credential.oauth2.access_token
|
||||
== 'some_access_token'
|
||||
)
|
||||
|
||||
# Check exchanged_auth_credential fields
|
||||
assert (
|
||||
called_config.exchanged_auth_credential.oauth2.client_id
|
||||
== 'expected_client_id'
|
||||
)
|
||||
assert (
|
||||
called_config.exchanged_auth_credential.oauth2.client_secret
|
||||
== 'expected_client_secret'
|
||||
)
|
||||
assert (
|
||||
called_config.exchanged_auth_credential.oauth2.redirect_uri
|
||||
== 'expected_redirect_uri'
|
||||
)
|
||||
assert (
|
||||
called_config.exchanged_auth_credential.oauth2.code_verifier
|
||||
== 'expected_code_verifier'
|
||||
)
|
||||
assert (
|
||||
called_config.exchanged_auth_credential.oauth2.code_challenge_method
|
||||
== 'S256'
|
||||
)
|
||||
assert (
|
||||
called_config.exchanged_auth_credential.oauth2.token_endpoint_auth_method
|
||||
== 'client_secret_post'
|
||||
)
|
||||
assert (
|
||||
called_config.exchanged_auth_credential.oauth2.access_token
|
||||
== 'some_exchanged_token'
|
||||
)
|
||||
|
||||
def test_merge_credential_oauth2_fields_when_target_oauth2_is_none(self):
|
||||
"""Test merging fields into a target credential where target.oauth2 is None."""
|
||||
from google.adk.auth.auth_preprocessor import _merge_credential_oauth2_fields
|
||||
|
||||
target = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
oauth2=None,
|
||||
)
|
||||
source = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id='expected_client_id',
|
||||
client_secret='expected_client_secret',
|
||||
),
|
||||
)
|
||||
|
||||
merged = _merge_credential_oauth2_fields(target, source)
|
||||
assert merged is not None
|
||||
assert merged.oauth2 is not None
|
||||
assert merged.oauth2.client_id == 'expected_client_id'
|
||||
assert merged.oauth2.client_secret == 'expected_client_secret'
|
||||
|
||||
@@ -406,9 +406,12 @@ class TestBuildAuthRequestEvent:
|
||||
self, mock_invocation_context
|
||||
):
|
||||
"""Test that multiple auth requests create multiple function call parts."""
|
||||
config1 = create_oauth2_auth_config()
|
||||
config2 = create_oauth2_auth_config()
|
||||
config2.credential_key = "different_key"
|
||||
auth_requests = {
|
||||
"call_1": create_oauth2_auth_config(),
|
||||
"call_2": create_oauth2_auth_config(),
|
||||
"call_1": config1,
|
||||
"call_2": config2,
|
||||
}
|
||||
|
||||
event = build_auth_request_event(mock_invocation_context, auth_requests)
|
||||
@@ -419,6 +422,27 @@ class TestBuildAuthRequestEvent:
|
||||
}
|
||||
assert function_call_ids == {"call_1", "call_2"}
|
||||
|
||||
def test_duplicate_auth_requests_are_deduplicated(
|
||||
self, mock_invocation_context
|
||||
):
|
||||
"""Test that auth requests with the same credential key are deduplicated."""
|
||||
config1 = create_oauth2_auth_config()
|
||||
config2 = create_oauth2_auth_config()
|
||||
# Ensure they have the same credential key
|
||||
assert config1.credential_key == config2.credential_key
|
||||
|
||||
auth_requests = {
|
||||
"call_1": config1,
|
||||
"call_2": config2,
|
||||
}
|
||||
|
||||
event = build_auth_request_event(mock_invocation_context, auth_requests)
|
||||
|
||||
assert len(event.content.parts) == 1
|
||||
fc = event.content.parts[0].function_call
|
||||
assert fc.name == REQUEST_EUC_FUNCTION_CALL_NAME
|
||||
assert fc.args["functionCallId"] == "call_1"
|
||||
|
||||
def test_always_adds_long_running_tool_ids(self, mock_invocation_context):
|
||||
"""Test that long_running_tool_ids is always set."""
|
||||
auth_requests = {"call_123": create_oauth2_auth_config()}
|
||||
|
||||
@@ -2642,8 +2642,12 @@ async def test_generate_auth_event_emits_one_long_running_call_per_request():
|
||||
function_response_event = _tool_response_event(
|
||||
invocation_context,
|
||||
{
|
||||
'orig_call_1': AuthConfig(auth_scheme=HTTPBearer()),
|
||||
'orig_call_2': AuthConfig(auth_scheme=HTTPBearer()),
|
||||
'orig_call_1': AuthConfig(
|
||||
auth_scheme=HTTPBearer(), credential_key='key1'
|
||||
),
|
||||
'orig_call_2': AuthConfig(
|
||||
auth_scheme=HTTPBearer(), credential_key='key2'
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
@@ -2664,6 +2668,34 @@ async def test_generate_auth_event_emits_one_long_running_call_per_request():
|
||||
] == ['orig_call_1', 'orig_call_2']
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_auth_event_deduplicates_requests():
|
||||
"""Duplicate requests for the same credential only emit one client-side call."""
|
||||
_, invocation_context = await _auth_invocation_context()
|
||||
function_response_event = _tool_response_event(
|
||||
invocation_context,
|
||||
{
|
||||
'orig_call_1': AuthConfig(
|
||||
auth_scheme=HTTPBearer(), credential_key='key1'
|
||||
),
|
||||
'orig_call_2': AuthConfig(
|
||||
auth_scheme=HTTPBearer(), credential_key='key1'
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
auth_event = generate_auth_event(invocation_context, function_response_event)
|
||||
|
||||
assert auth_event is not None
|
||||
calls = auth_event.get_function_calls()
|
||||
assert [call.name for call in calls] == [REQUEST_EUC_FUNCTION_CALL_NAME]
|
||||
assert len(calls) == 1
|
||||
assert (
|
||||
AuthToolArguments.model_validate(calls[0].args).function_call_id
|
||||
== 'orig_call_1'
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_auth_event_mirrors_the_tool_response_role():
|
||||
"""The auth request keeps the role of the tool response it came from."""
|
||||
|
||||
@@ -415,3 +415,40 @@ def test_legacy_credential_key_is_stable_across_redirect_uri():
|
||||
assert store._get_legacy_credential_key(
|
||||
scheme, credential_local
|
||||
) == store._get_legacy_credential_key(scheme, credential_deployed)
|
||||
|
||||
|
||||
def test_legacy_credential_migration(
|
||||
openid_connect_scheme, openid_connect_credential
|
||||
):
|
||||
"""Test that credentials stored under legacy keys are migrated to new keys."""
|
||||
tool_context = create_mock_tool_context()
|
||||
store = ToolContextCredentialStore(tool_context=tool_context)
|
||||
|
||||
legacy_key = store._get_legacy_credential_key(
|
||||
openid_connect_scheme, openid_connect_credential
|
||||
)
|
||||
new_key = store.get_credential_key(
|
||||
openid_connect_scheme, openid_connect_credential
|
||||
)
|
||||
assert legacy_key != new_key
|
||||
|
||||
legacy_credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.HTTP,
|
||||
http=HttpAuth(
|
||||
scheme='bearer',
|
||||
credentials=HttpCredentials(token='legacy_token'),
|
||||
),
|
||||
)
|
||||
store.store_credential(legacy_key, legacy_credential)
|
||||
|
||||
assert new_key not in tool_context.state
|
||||
|
||||
retrieved = store.get_credential(
|
||||
openid_connect_scheme, openid_connect_credential
|
||||
)
|
||||
|
||||
assert retrieved == legacy_credential
|
||||
assert new_key in tool_context.state
|
||||
assert tool_context.state[new_key] == legacy_credential.model_dump(
|
||||
exclude_none=True
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user