diff --git a/src/google/adk/integrations/agent_identity/_iam_connector_credentials_provider.py b/src/google/adk/integrations/agent_identity/_iam_connector_credentials_provider.py new file mode 100644 index 00000000..9d8e702f --- /dev/null +++ b/src/google/adk/integrations/agent_identity/_iam_connector_credentials_provider.py @@ -0,0 +1,272 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import asyncio +import logging +import os +import time + +from google.adk.agents.callback_context import CallbackContext +from google.adk.auth.auth_credential import AuthCredential +from google.adk.auth.auth_credential import AuthCredentialTypes +from google.adk.auth.auth_credential import HttpAuth +from google.adk.auth.auth_credential import HttpCredentials +from google.adk.auth.auth_credential import OAuth2Auth +from google.adk.flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME +from google.api_core.client_options import ClientOptions + +try: + from google.cloud.iamconnectorcredentials_v1alpha import IAMConnectorCredentialsServiceClient as Client + from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsMetadata + from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsRequest + from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsResponse +except ImportError as e: + raise ImportError( + "Missing required dependencies for Agent Identity Auth Manager. " + 'Please install with: pip install "google-adk[agent-identity]"' + ) from e +from google.longrunning.operations_pb2 import Operation + +from .gcp_auth_provider_scheme import GcpAuthProviderScheme + +# Notes on the current IAM Connector Credentials service implementation: +# 1. The service does not yet support LROs, so even though the +# retrieve_credentials method returns an Operation object, the methods like +# operation.done() and operation.result() will not work yet. +# 2. For API key flows, the returned Operation contains the credentials. +# 3. For 2-legged OAuth flows, the returned Operation contains pending status, +# client needs to retry the request until response with credentials is +# returned or timeout occurs. +# 4. For 3-legged OAuth flows, the returned Operation contains consent pending +# status along with the authorization URI. + +# TODO: Catch specific exceptions instead of generic ones. + +logger = logging.getLogger("google_adk." + __name__) + +NON_INTERACTIVE_TOKEN_POLL_INTERVAL_SEC: float = 1.0 +NON_INTERACTIVE_TOKEN_POLL_TIMEOUT_SEC: float = 10.0 + + +def _construct_auth_credential( + response: RetrieveCredentialsResponse, +) -> AuthCredential: + """Constructs a simplified HTTP auth credential from the header-token tuple returned by the upstream service.""" + if not response.header or not response.token: + raise ValueError( + "Received either empty header or token from IAM Connector Credentials" + " service." + ) + + header_name, _, header_value = response.header.partition(":") + if ( + header_name.strip().lower() == "authorization" + and header_value.strip().lower().startswith("bearer") + ): + return AuthCredential( + auth_type=AuthCredentialTypes.HTTP, + http=HttpAuth( + scheme="Bearer", + credentials=HttpCredentials(token=response.token), + ), + ) + + # Handle custom header. + return AuthCredential( + auth_type=AuthCredentialTypes.HTTP, + http=HttpAuth( + # For custom headers, scheme and credentials fields are not used. + scheme="", + credentials=HttpCredentials(), + additional_headers={ + response.header: response.token, + "X-GOOG-API-KEY": response.token, + }, + ), + ) + + +class _IamConnectorCredentialsProvider: + """Implementation for auth provider using IAM Connector credentials service.""" + + _client: Client | None = None + + def __init__(self, client: Client | None = None): + self._client = client + + def _get_client(self) -> Client: + """Lazy loads the client to avoid unnecessary setup on startup.""" + if self._client is None: + client_options = None + if host := os.environ.get("IAM_CONNECTOR_CREDENTIALS_TARGET_HOST"): + client_options = ClientOptions(api_endpoint=host) + self._client = Client(client_options=client_options, transport="rest") + return self._client + + async def _retrieve_credentials( + self, + user_id: str, + auth_scheme: GcpAuthProviderScheme, + ) -> Operation: + request = RetrieveCredentialsRequest( + connector=auth_scheme.name, + user_id=user_id, + scopes=auth_scheme.scopes, + continue_uri=auth_scheme.continue_uri or "", + force_refresh=False, + ) + # TODO: Use async client once available. Temporarily using threading to + # prevent blocking the event loop. + operation = await asyncio.to_thread( + self._get_client().retrieve_credentials, request + ) + return operation.operation + + def _unpack_operation( + self, operation: Operation + ) -> tuple[ + RetrieveCredentialsResponse | None, RetrieveCredentialsMetadata | None + ]: + """Deserializes the response and metadata from the operation.""" + response = None + metadata = None + if operation.response: + response = RetrieveCredentialsResponse.deserialize( + operation.response.value + ) + if operation.metadata: + metadata = RetrieveCredentialsMetadata.deserialize( + operation.metadata.value + ) + return response, metadata + + async def _poll_credentials( + self, user_id: str, auth_scheme: GcpAuthProviderScheme, timeout: float + ) -> Operation: + end_time = time.time() + timeout + while time.time() < end_time: + operation = await self._retrieve_credentials(user_id, auth_scheme) + if operation.done: + return operation + await asyncio.sleep(NON_INTERACTIVE_TOKEN_POLL_INTERVAL_SEC) + raise TimeoutError("Timeout waiting for credentials.") + + @staticmethod + def _is_consent_completed(context: CallbackContext) -> bool: + """Checks if the user consent flow is completed for the current function call.""" + if not context.function_call_id: + return False + + if not context.session: + return False + + events = context.session.events + target_tool_call_id = context.function_call_id + + # Find all relevant function calls and responses + euc_calls = {} + euc_responses = {} + + for event in events: + for call in event.get_function_calls(): + if call.name == REQUEST_EUC_FUNCTION_CALL_NAME: + euc_calls[call.id] = call + for response in event.get_function_responses(): + if response.name == REQUEST_EUC_FUNCTION_CALL_NAME: + euc_responses[response.id] = response + + # Check for a response that matches a call for the current tool invocation + for call_id, _ in euc_responses.items(): + if call_id in euc_calls: + call = euc_calls[call_id] + if call.args and call.args.get("functionCallId") == target_tool_call_id: + return True + return False + + async def get_auth_credential( + self, + auth_scheme: GcpAuthProviderScheme, + context: CallbackContext | None = None, + ) -> AuthCredential: + """Retrieves credentials using the IAM Connector Credentials service. + + Args: + auth_scheme: The GcpAuthProviderScheme. + context: Optional context for the callback. + + Returns: + An AuthCredential instance. + + Raises: + RuntimeError: If credential retrieval or polling fails. + """ + + if context is None or context.user_id is None: + raise ValueError( + "GcpAuthProvider requires a context with a valid user_id." + ) + + user_id = context.user_id + + try: + operation = await self._retrieve_credentials(user_id, auth_scheme) + except Exception as e: + raise RuntimeError( + f"Failed to retrieve credential for user '{user_id}' on connector" + f" '{auth_scheme.name}'." + ) from e + + response, metadata = self._unpack_operation(operation) + + if operation.HasField("error"): + raise RuntimeError(f"Operation failed: {operation.error.message}") + + if operation.done: + logger.debug("Auth credential obtained immediately.") + return _construct_auth_credential(response) + + if metadata and metadata.consent_pending: + # Get 2-legged OAuth token. Allow enough time for token exchange. + try: + operation = await self._poll_credentials( + user_id, + auth_scheme, + timeout=NON_INTERACTIVE_TOKEN_POLL_TIMEOUT_SEC, + ) + if operation.HasField("error"): + raise RuntimeError(f"Operation failed: {operation.error.message}") + if operation.done: + logger.debug("Auth credential obtained after polling.") + response, _ = self._unpack_operation(operation) + return _construct_auth_credential(response) + except Exception as e: + raise RuntimeError( + f"Failed to retrieve credential for user '{user_id}' on connector" + f" '{auth_scheme.name}'." + ) from e + + if metadata is not None and metadata.uri_consent_required: + if self._is_consent_completed(context): + raise RuntimeError("Failed to retrieve consent based credential.") + + # Return AuthCredential with only auth_uri to trigger user consent flow. + return AuthCredential( + auth_type=AuthCredentialTypes.OAUTH2, + oauth2=OAuth2Auth( + auth_uri=metadata.uri_consent_required.authorization_uri, + nonce=metadata.uri_consent_required.consent_nonce, + ), + ) diff --git a/src/google/adk/integrations/agent_identity/gcp_auth_provider.py b/src/google/adk/integrations/agent_identity/gcp_auth_provider.py index 355faff0..eaffc5a4 100644 --- a/src/google/adk/integrations/agent_identity/gcp_auth_provider.py +++ b/src/google/adk/integrations/agent_identity/gcp_auth_provider.py @@ -12,198 +12,32 @@ # See the License for the specific language governing permissions and # limitations under the License. +"""Authentication provider using Google Cloud Agent Identity Credentials service.""" + from __future__ import annotations -import asyncio -import logging -import os -import time from google.adk.agents.callback_context import CallbackContext from google.adk.auth.auth_credential import AuthCredential -from google.adk.auth.auth_credential import AuthCredentialTypes -from google.adk.auth.auth_credential import HttpAuth -from google.adk.auth.auth_credential import HttpCredentials -from google.adk.auth.auth_credential import OAuth2Auth from google.adk.auth.auth_tool import AuthConfig from google.adk.auth.base_auth_provider import BaseAuthProvider -from google.adk.flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME -from google.api_core.client_options import ClientOptions - -try: - from google.cloud.iamconnectorcredentials_v1alpha import IAMConnectorCredentialsServiceClient as Client - from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsMetadata - from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsRequest - from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsResponse -except ImportError as e: - raise ImportError( - "Missing required dependencies for Agent Identity Auth Manager. " - 'Please install with: pip install "google-adk[agent-identity]"' - ) from e -from google.longrunning.operations_pb2 import Operation from typing_extensions import override +from ._iam_connector_credentials_provider import _IamConnectorCredentialsProvider from .gcp_auth_provider_scheme import GcpAuthProviderScheme -# Notes on the current Agent Identity Credentials service implementation: -# 1. The service does not yet support LROs, so even though the -# retrieve_credentials method returns an Operation object, the methods like -# operation.done() and operation.result() will not work yet. -# 2. For API key flows, the returned Operation contains the credentials. -# 3. For 2-legged OAuth flows, the returned Operation contains pending status, -# client needs to retry the request until response with credentials is -# returned or timeout occurs. -# 4. For 3-legged OAuth flows, the returned Operation contains consent pending -# status along with the authorization URI. - -# TODO: Catch specific exceptions instead of generic ones. - -logger = logging.getLogger("google_adk." + __name__) - -NON_INTERACTIVE_TOKEN_POLL_INTERVAL_SEC: float = 1.0 -NON_INTERACTIVE_TOKEN_POLL_TIMEOUT_SEC: float = 10.0 - - -def _construct_auth_credential( - response: RetrieveCredentialsResponse, -) -> AuthCredential: - """Constructs a simplified HTTP auth credential from the header-token tuple returned by the upstream service.""" - if not response.header or not response.token: - raise ValueError( - "Received either empty header or token from Agent Identity Credentials" - " service." - ) - - header_name, _, header_value = response.header.partition(":") - if ( - header_name.strip().lower() == "authorization" - and header_value.strip().lower().startswith("bearer") - ): - return AuthCredential( - auth_type=AuthCredentialTypes.HTTP, - http=HttpAuth( - scheme="Bearer", - credentials=HttpCredentials(token=response.token), - ), - ) - - # Handle custom header. - return AuthCredential( - auth_type=AuthCredentialTypes.HTTP, - http=HttpAuth( - # For custom headers, scheme and credentials fields are not used. - scheme="", - credentials=HttpCredentials(), - additional_headers={ - response.header: response.token, - "X-GOOG-API-KEY": response.token, - }, - ), - ) - class GcpAuthProvider(BaseAuthProvider): """An auth provider that uses the Agent Identity Credentials service to generate access tokens.""" - _client: Client | None = None - - def __init__(self, client: Client | None = None): - self._client = client + def __init__(self): + self._iam_connector_provider = _IamConnectorCredentialsProvider() @property @override def supported_auth_schemes(self) -> tuple[type[GcpAuthProviderScheme], ...]: return (GcpAuthProviderScheme,) - def _get_client(self) -> Client: - """Lazy loads the client to avoid unnecessary setup on startup.""" - if self._client is None: - client_options = None - if host := os.environ.get("IAM_CONNECTOR_CREDENTIALS_TARGET_HOST"): - client_options = ClientOptions(api_endpoint=host) - self._client = Client(client_options=client_options, transport="rest") - return self._client - - async def _retrieve_credentials( - self, - user_id: str, - auth_scheme: GcpAuthProviderScheme, - ) -> Operation: - request = RetrieveCredentialsRequest( - connector=auth_scheme.name, - user_id=user_id, - scopes=auth_scheme.scopes, - continue_uri=auth_scheme.continue_uri or "", - force_refresh=False, - ) - # TODO: Use async client once available. Temporarily using threading to - # prevent blocking the event loop. - operation = await asyncio.to_thread( - self._get_client().retrieve_credentials, request - ) - return operation.operation - - def _unpack_operation( - self, operation: Operation - ) -> tuple[ - RetrieveCredentialsResponse | None, RetrieveCredentialsMetadata | None - ]: - """Deserializes the response and metadata from the operation.""" - response = None - metadata = None - if operation.response: - response = RetrieveCredentialsResponse.deserialize( - operation.response.value - ) - if operation.metadata: - metadata = RetrieveCredentialsMetadata.deserialize( - operation.metadata.value - ) - return response, metadata - - async def _poll_credentials( - self, user_id: str, auth_scheme: GcpAuthProviderScheme, timeout: float - ) -> Operation: - end_time = time.time() + timeout - while time.time() < end_time: - operation = await self._retrieve_credentials(user_id, auth_scheme) - if operation.done: - return operation - await asyncio.sleep(NON_INTERACTIVE_TOKEN_POLL_INTERVAL_SEC) - raise TimeoutError("Timeout waiting for credentials.") - - @staticmethod - def _is_consent_completed(context: CallbackContext) -> bool: - """Checks if the user consent flow is completed for the current function call.""" - if not context.function_call_id: - return False - - if not context.session: - return False - - events = context.session.events - target_tool_call_id = context.function_call_id - - # Find all relevant function calls and responses - euc_calls = {} - euc_responses = {} - - for event in events: - for call in event.get_function_calls(): - if call.name == REQUEST_EUC_FUNCTION_CALL_NAME: - euc_calls[call.id] = call - for response in event.get_function_responses(): - if response.name == REQUEST_EUC_FUNCTION_CALL_NAME: - euc_responses[response.id] = response - - # Check for a response that matches a call for the current tool invocation - for call_id, _ in euc_responses.items(): - if call_id in euc_calls: - call = euc_calls[call_id] - if call.args and call.args.get("functionCallId") == target_tool_call_id: - return True - return False - @override async def get_auth_credential( self, @@ -221,68 +55,13 @@ class GcpAuthProvider(BaseAuthProvider): Raises: ValueError: If auth_scheme is not a GcpAuthProviderScheme. - RuntimeError: If credential retrieval or polling fails. """ - auth_scheme = auth_config.auth_scheme if not isinstance(auth_scheme, GcpAuthProviderScheme): raise ValueError( f"Expected GcpAuthProviderScheme, got {type(auth_scheme)}" ) - if context is None or context.user_id is None: - raise ValueError( - "GcpAuthProvider requires a context with a valid user_id." - ) - - user_id = context.user_id - - try: - operation = await self._retrieve_credentials(user_id, auth_scheme) - except Exception as e: - raise RuntimeError( - f"Failed to retrieve credential for user '{user_id}' on connector" - f" '{auth_scheme.name}'." - ) from e - - response, metadata = self._unpack_operation(operation) - - if operation.HasField("error"): - raise RuntimeError(f"Operation failed: {operation.error.message}") - - if operation.done: - logger.debug("Auth credential obtained immediately.") - return _construct_auth_credential(response) - - if metadata and metadata.consent_pending: - # Get 2-legged OAuth token. Allow enough time for token exchange. - try: - operation = await self._poll_credentials( - user_id, - auth_scheme, - timeout=NON_INTERACTIVE_TOKEN_POLL_TIMEOUT_SEC, - ) - if operation.HasField("error"): - raise RuntimeError(f"Operation failed: {operation.error.message}") - if operation.done: - logger.debug("Auth credential obtained after polling.") - response, _ = self._unpack_operation(operation) - return _construct_auth_credential(response) - except Exception as e: - raise RuntimeError( - f"Failed to retrieve credential for user '{user_id}' on connector" - f" '{auth_scheme.name}'." - ) from e - - if metadata is not None and metadata.uri_consent_required: - if self._is_consent_completed(context): - raise RuntimeError("Failed to retrieve consent based credential.") - - # Return AuthCredential with only auth_uri to trigger user consent flow. - return AuthCredential( - auth_type=AuthCredentialTypes.OAUTH2, - oauth2=OAuth2Auth( - auth_uri=metadata.uri_consent_required.authorization_uri, - nonce=metadata.uri_consent_required.consent_nonce, - ), - ) + return await self._iam_connector_provider.get_auth_credential( + auth_scheme=auth_scheme, context=context + ) diff --git a/tests/integration/integrations/agent_identity/test_2lo_flow.py b/tests/integration/integrations/agent_identity/test_2lo_flow.py index b9431cfc..45c5abb5 100644 --- a/tests/integration/integrations/agent_identity/test_2lo_flow.py +++ b/tests/integration/integrations/agent_identity/test_2lo_flow.py @@ -22,7 +22,7 @@ from google.adk import Agent from google.adk import Runner from google.adk.auth.auth_tool import AuthConfig from google.adk.auth.credential_manager import CredentialManager -from google.adk.integrations.agent_identity import gcp_auth_provider +from google.adk.integrations.agent_identity import _iam_connector_credentials_provider from google.adk.integrations.agent_identity import GcpAuthProvider from google.adk.integrations.agent_identity import GcpAuthProviderScheme from google.adk.sessions.in_memory_session_service import InMemorySessionService @@ -103,7 +103,7 @@ async def test_gcp_agent_identity_2lo_gets_token() -> None: # 1. Setup mocked GCP Client to return the fake Bearer token with mock.patch.object( - gcp_auth_provider, + _iam_connector_credentials_provider, "Client", autospec=True, ) as mock_client_cls: @@ -209,7 +209,7 @@ async def test_gcp_agent_identity_2lo_sends_authorization_header_to_mcp_session( mock_operation = _DummyOperation() with mock.patch.object( - gcp_auth_provider, "Client", autospec=True + _iam_connector_credentials_provider, "Client", autospec=True ) as mock_gcp: mock_gcp.return_value.retrieve_credentials.return_value = mock_operation diff --git a/tests/integration/integrations/agent_identity/test_3lo_flow.py b/tests/integration/integrations/agent_identity/test_3lo_flow.py index 767d51a2..74dbcb54 100644 --- a/tests/integration/integrations/agent_identity/test_3lo_flow.py +++ b/tests/integration/integrations/agent_identity/test_3lo_flow.py @@ -22,7 +22,7 @@ from google.adk import Agent from google.adk import Runner from google.adk.auth.auth_tool import AuthConfig from google.adk.auth.credential_manager import CredentialManager -from google.adk.integrations.agent_identity import gcp_auth_provider +from google.adk.integrations.agent_identity import _iam_connector_credentials_provider from google.adk.integrations.agent_identity import GcpAuthProvider from google.adk.integrations.agent_identity import GcpAuthProviderScheme from google.adk.sessions.in_memory_session_service import InMemorySessionService @@ -146,7 +146,7 @@ async def test_gcp_agent_identity_3lo_user_consent_flow() -> None: mock_gcp_client = MockGcpClient() with mock.patch.object( - gcp_auth_provider, + _iam_connector_credentials_provider, "Client", autospec=True, ) as mock_client_cls: diff --git a/tests/unittests/integrations/agent_identity/test_gcp_auth_provider.py b/tests/unittests/integrations/agent_identity/test_gcp_auth_provider.py index dddd8b7d..feedf879 100644 --- a/tests/unittests/integrations/agent_identity/test_gcp_auth_provider.py +++ b/tests/unittests/integrations/agent_identity/test_gcp_auth_provider.py @@ -12,40 +12,17 @@ # See the License for the specific language governing permissions and # limitations under the License. +from unittest.mock import AsyncMock from unittest.mock import Mock from unittest.mock import patch -import pytest - -pytest.importorskip( - "google.cloud.iamconnectorcredentials_v1alpha", - reason="Requires google-cloud-iamconnectorcredentials", -) - from google.adk.agents.callback_context import CallbackContext -from google.adk.auth.auth_credential import AuthCredentialTypes +from google.adk.auth.auth_credential import AuthCredential from google.adk.auth.auth_tool import AuthConfig -from google.adk.auth.auth_tool import AuthToolArguments -from google.adk.flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME -from google.adk.integrations.agent_identity import gcp_auth_provider from google.adk.integrations.agent_identity import GcpAuthProvider from google.adk.integrations.agent_identity import GcpAuthProviderScheme -from google.adk.sessions.session import Session -from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsMetadata -from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsResponse -from google.longrunning.operations_pb2 import Operation -from google.protobuf.any_pb2 import Any -from google.rpc.status_pb2 import Status - - -@pytest.fixture -def mock_client(): - return Mock(spec=gcp_auth_provider.Client) - - -@pytest.fixture -def provider(mock_client): - return GcpAuthProvider(client=mock_client) +from google.adk.integrations.agent_identity._iam_connector_credentials_provider import _IamConnectorCredentialsProvider +import pytest @pytest.fixture @@ -58,420 +35,43 @@ def auth_config(): return Mock(spec=AuthConfig, auth_scheme=scheme) -@pytest.fixture -def mock_operation(mocker, mock_client): - op = Operation(done=True) - - class DummyCall: - - def __init__(self, operation): - self.operation = operation - - mock_client.retrieve_credentials.return_value = DummyCall(op) - return op - - @pytest.fixture def context(): context = Mock(spec=CallbackContext) context.user_id = "user" - context.function_call_id = "call_123" - session = Mock(spec=Session) - session.events = [] - context.session = session - return context +@pytest.fixture +def provider(): + return GcpAuthProvider() -@patch.dict(gcp_auth_provider.os.environ, clear=True) -@patch.object(gcp_auth_provider, "Client") -def test_get_client_uses_rest_transport(mock_client_class): +def test_supported_auth_schemes(provider): + """Verify the provider supports the correct auth scheme.""" + assert GcpAuthProviderScheme in provider.supported_auth_schemes + + +@patch("google.adk.integrations.agent_identity.gcp_auth_provider._IamConnectorCredentialsProvider") +async def test_gcp_auth_provider_delegates_get_auth_credential(mock_provider_class, auth_config, context): + """Test that get_auth_credential delegates to the internal provider.""" provider = GcpAuthProvider() - provider._get_client() - mock_client_class.assert_called_once() - _, kwargs = mock_client_class.call_args - assert kwargs.get("transport") == "rest" + mock_credential = Mock(spec=AuthCredential) + mock_provider_instance = mock_provider_class.return_value + mock_provider_instance.get_auth_credential = AsyncMock(return_value=mock_credential) + result = await provider.get_auth_credential(auth_config, context) -@patch.dict( - gcp_auth_provider.os.environ, - {"IAM_CONNECTOR_CREDENTIALS_TARGET_HOST": "some-host"}, -) -@patch.object(gcp_auth_provider, "Client") -@patch.object(gcp_auth_provider, "ClientOptions") -def test_get_client_with_env_var(mock_client_options_class, mock_client_class): - provider = GcpAuthProvider() - client = provider._get_client() - - assert client == mock_client_class.return_value - mock_client_options_class.assert_called_once_with(api_endpoint="some-host") - mock_client_class.assert_called_once_with( - client_options=mock_client_options_class.return_value, transport="rest" + assert result == mock_credential + mock_provider_instance.get_auth_credential.assert_awaited_once_with( + auth_scheme=auth_config.auth_scheme, context=context ) -# ============================================================================== -# Non-interactive auth flows (API key and 2-legged OAuth) -# ============================================================================== - - -async def test_get_auth_credential_raises_error_for_invalid_auth_scheme( - provider, context -): +async def test_get_auth_credential_raises_error_for_invalid_auth_scheme(context): """Test get_auth_credential raises ValueError for invalid auth scheme.""" + provider = GcpAuthProvider() invalid_auth_config = Mock(spec=AuthConfig) invalid_auth_config.auth_scheme = Mock() # Not GcpAuthProviderScheme with pytest.raises(ValueError, match="Expected GcpAuthProviderScheme, got"): await provider.get_auth_credential(invalid_auth_config, context) - - -async def test_get_auth_credential_raises_error_if_context_is_missing( - provider, auth_config -): - """Test get_auth_credential raises ValueError if context is missing.""" - with pytest.raises( - ValueError, - match="GcpAuthProvider requires a context with a valid user_id", - ): - await provider.get_auth_credential(auth_config, context=None) - - -async def test_get_auth_credential_raises_error_if_user_id_is_missing( - provider, auth_config -): - """Test get_auth_credential raises ValueError if user_id is missing.""" - context = Mock(spec=CallbackContext) - context.user_id = None - with pytest.raises( - ValueError, - match="GcpAuthProvider requires a context with a valid user_id", - ): - await provider.get_auth_credential(auth_config, context=context) - - -async def test_get_auth_credential_returns_credential_if_available_immediately( - mock_client, - mock_operation, - auth_config, - context, - provider, -): - """Test get_auth_credential returns credential if available immediately.""" - mock_credential = RetrieveCredentialsResponse( - header="Authorization: Bearer", token="test-token" - ) - mock_operation.response.value = RetrieveCredentialsResponse.serialize( - mock_credential - ) - - auth_credential = await provider.get_auth_credential(auth_config, context) - - assert auth_credential.auth_type == AuthCredentialTypes.HTTP - assert auth_credential.http.scheme == "Bearer" - assert auth_credential.http.credentials.token == "test-token" - mock_client.retrieve_credentials.assert_called_once() - - -async def test_get_auth_credential_raises_error_if_upstream_returns_empty_header( - mock_operation, - auth_config, - context, - provider, -): - """Test get_auth_credential raises RuntimeError for empty header.""" - mock_credential = RetrieveCredentialsResponse(header="", token="test-token") - mock_operation.response.value = RetrieveCredentialsResponse.serialize( - mock_credential - ) - - with pytest.raises( - ValueError, - match=( - "Received either empty header or token from Agent Identity" - " Credentials service." - ), - ): - await provider.get_auth_credential(auth_config, context) - - -async def test_get_auth_credential_raises_error_if_upstream_returns_empty_token( - mock_operation, - auth_config, - context, - provider, -): - """Test get_auth_credential raises RuntimeError for empty token.""" - mock_credential = RetrieveCredentialsResponse( - header="Authorization: Bearer", token="" - ) - mock_operation.response.value = RetrieveCredentialsResponse.serialize( - mock_credential - ) - - with pytest.raises( - ValueError, - match=( - "Received either empty header or token from Agent Identity" - " Credentials service." - ), - ): - await provider.get_auth_credential(auth_config, context) - - -async def test_get_auth_credential_returns_credential_if_upstream_returns_custom_header( - mock_operation, - auth_config, - context, - provider, -): - """Test get_auth_credential returns valid credential for custom header and sets X-GOOG-API-KEY header.""" - mock_credential = RetrieveCredentialsResponse( - header="some-x-api-key", token="test-token" - ) - mock_operation.response.value = RetrieveCredentialsResponse.serialize( - mock_credential - ) - - auth_credential = await provider.get_auth_credential(auth_config, context) - - assert auth_credential.auth_type == AuthCredentialTypes.HTTP - assert not auth_credential.http.scheme - assert auth_credential.http.credentials.token is None - assert auth_credential.http.additional_headers == { - "some-x-api-key": "test-token", - "X-GOOG-API-KEY": "test-token", - } - - -async def test_get_auth_credential_raises_error_if_upstream_operation_errors( - mock_operation, auth_config, context, provider -): - """Test get_auth_credential raises RuntimeError for failed operations.""" - mock_operation.error.message = "OAuth server error" - mock_operation.done = False - - with pytest.raises( - RuntimeError, match="Operation failed: OAuth server error" - ): - await provider.get_auth_credential(auth_config, context) - - -async def test_get_auth_credential_raises_error_if_upstream_call_fails( - mock_client, auth_config, context, provider -): - """Test get_auth_credential raises RuntimeError for failed calls.""" - mock_client.retrieve_credentials.side_effect = Exception( - "API Quota Exhausted" - ) - - with pytest.raises( - RuntimeError, - match="Failed to retrieve credential for user 'user' on connector", - ) as exc_info: - await provider.get_auth_credential(auth_config, context) - - # Assert that the original Exception is the chained cause! - assert str(exc_info.value.__cause__) == "API Quota Exhausted" - - -@patch.object(gcp_auth_provider.time, "time") -async def test_get_auth_credential_raises_error_if_polling_times_out( - mock_time, - mock_operation, - auth_config, - context, - provider, -): - """Test get_auth_credential raises RuntimeError if polling times out.""" - - # Force the operation into the polling loop state - meta_pb = RetrieveCredentialsMetadata.pb()() - meta_pb.consent_pending.SetInParent() - meta = RetrieveCredentialsMetadata.deserialize(meta_pb.SerializeToString()) - mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta) - - # First call sets start_time=0.0, second call checks time > timeout - # (20.0 > 10.0) - mock_time.side_effect = [0.0, 20.0] - - mock_metadata = Mock(spec=RetrieveCredentialsMetadata) - mock_metadata.consent_pending = True - mock_metadata.uri_consent_required = False - mock_operation.done = True - mock_operation.ClearField("error") - mock_client = Mock(spec=gcp_auth_provider.Client) - mock_client.retrieve_credentials.side_effect = Exception( - "Timeout waiting for credentials." - ) - provider._client = mock_client - - with pytest.raises( - RuntimeError, - match="Failed to retrieve credential for user 'user' on connector", - ) as exc_info: - await provider.get_auth_credential(auth_config, context) - - assert "Timeout waiting for credentials." in str(exc_info.value.__cause__) - - -# ============================================================================== -# Interactive Auth Flows (3-legged OAuth for User Consents) -# ============================================================================== - - -async def test_get_auth_credential_initiates_user_consent( - mock_operation, auth_config, context, provider -): - # Explicitly set the mock behavior for this test - expected_uri = "https://example.com/auth" - expected_nonce = "sample-nonce-123" - meta = RetrieveCredentialsMetadata({ - "uri_consent_required": { - "authorization_uri": expected_uri, - "consent_nonce": expected_nonce, - } - }) - mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta) - mock_operation.done = False - # Assert that there is no prior user consent completion event - assert not context.session.events - - credential = await provider.get_auth_credential(auth_config, context) - - assert credential is not None - assert credential.auth_type == AuthCredentialTypes.OAUTH2 - assert credential.oauth2.auth_uri == expected_uri - assert credential.oauth2.nonce == expected_nonce - - -async def test_get_auth_credential_returns_fresh_auth_uri_for_repeated_requests( - mock_client, mock_operation, auth_config, context, provider -): - """Test that repeated calls fetch fresh auth URIs if consent is still pending.""" - # Arrange: Explicit initial URI - initial_uri = "https://example.com/auth" - initial_nonce = "initial-nonce-123" - meta1 = RetrieveCredentialsMetadata({ - "uri_consent_required": { - "authorization_uri": initial_uri, - "consent_nonce": initial_nonce, - } - }) - mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta1) - mock_operation.done = False - - credential1 = await provider.get_auth_credential(auth_config, context) - assert credential1.oauth2.auth_uri == initial_uri - assert credential1.oauth2.nonce == initial_nonce - - # Arrange: Explicit new URI for the second call - fresh_auth_uri = "https://example.com/auth_new" - fresh_nonce = "fresh-nonce-456" - meta2 = RetrieveCredentialsMetadata({ - "uri_consent_required": { - "authorization_uri": fresh_auth_uri, - "consent_nonce": fresh_nonce, - } - }) - mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta2) - - credential2 = await provider.get_auth_credential(auth_config, context) - - assert mock_client.retrieve_credentials.call_count == 2 - assert credential2.oauth2.auth_uri == fresh_auth_uri - assert credential2.oauth2.nonce == fresh_nonce - - -async def test_get_auth_credential_returns_token_if_consent_was_completed( - mock_operation, auth_config, context, provider -): - # Setup mock credential for successful credential retrieval - mock_credential = RetrieveCredentialsResponse( - header="Authorization: Bearer", token="test-token" - ) - mock_operation.response.value = RetrieveCredentialsResponse.serialize( - mock_credential - ) - - # Create mock events - # 1. FunctionCall event for adk_request_credential - function_call = Mock() - function_call.id = "auth-req-1" - function_call.name = REQUEST_EUC_FUNCTION_CALL_NAME - function_call.args = AuthToolArguments( - function_call_id="call-123", auth_config=auth_config - ).model_dump(by_alias=True, exclude_none=True) - - event1 = Mock() - event1.get_function_calls.return_value = [function_call] - event1.get_function_responses.return_value = [] - - # 2. FunctionResponse event for adk_request_credential - function_response = Mock() - function_response.id = "auth-req-1" - function_response.name = REQUEST_EUC_FUNCTION_CALL_NAME - - event2 = Mock() - event2.get_function_calls.return_value = [] - event2.get_function_responses.return_value = [function_response] - - # Setup tool context and event history (order of events matters) - context.session.events = [event1, event2] - context.function_call_id = "call-123" - - # Also set uri_consent_required to True-ish so it enters the check block - meta = RetrieveCredentialsMetadata( - uri_consent_required=RetrieveCredentialsMetadata.UriConsentRequired() - ) - mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta) - - # Execute - auth_credential = await provider.get_auth_credential(auth_config, context) - - # Verify - assert auth_credential is not None - assert auth_credential.auth_type == AuthCredentialTypes.HTTP - assert auth_credential.http.scheme == "Bearer" - assert auth_credential.http.credentials.token == "test-token" - - -async def test_get_auth_credential_raises_error_if_consent_canceled( - mock_operation, auth_config, context, provider -): - function_call = Mock() - function_call.id = "auth-req-1" - function_call.name = REQUEST_EUC_FUNCTION_CALL_NAME - function_call.args = AuthToolArguments( - function_call_id="call-123", auth_config=auth_config - ).model_dump(by_alias=True, exclude_none=True) - - event1 = Mock() - event1.get_function_calls.return_value = [function_call] - event1.get_function_responses.return_value = [] - - function_response = Mock() - function_response.id = "auth-req-1" - function_response.name = REQUEST_EUC_FUNCTION_CALL_NAME - - event2 = Mock() - event2.get_function_calls.return_value = [] - event2.get_function_responses.return_value = [function_response] - - context.session.events = [event1, event2] - context.function_call_id = "call-123" - - meta = RetrieveCredentialsMetadata({ - "uri_consent_required": { - "authorization_uri": "https://example.com/auth", - "consent_nonce": "sample-nonce", - } - }) - mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta) - mock_operation.done = False - - with pytest.raises( - RuntimeError, match="Failed to retrieve consent based credential." - ): - await provider.get_auth_credential(auth_config, context) diff --git a/tests/unittests/integrations/agent_identity/test_iam_connector_credentials_provider.py b/tests/unittests/integrations/agent_identity/test_iam_connector_credentials_provider.py new file mode 100644 index 00000000..a3333120 --- /dev/null +++ b/tests/unittests/integrations/agent_identity/test_iam_connector_credentials_provider.py @@ -0,0 +1,466 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from unittest.mock import Mock +from unittest.mock import patch + +import pytest + +pytest.importorskip( + "google.cloud.iamconnectorcredentials_v1alpha", + reason="Requires google-cloud-iamconnectorcredentials", +) + +from google.adk.agents.callback_context import CallbackContext +from google.adk.auth.auth_credential import AuthCredentialTypes +from google.adk.auth.auth_tool import AuthConfig +from google.adk.auth.auth_tool import AuthToolArguments +from google.adk.flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME +from google.adk.integrations.agent_identity import _iam_connector_credentials_provider +from google.adk.integrations.agent_identity import GcpAuthProviderScheme +from google.adk.integrations.agent_identity._iam_connector_credentials_provider import _IamConnectorCredentialsProvider +from google.adk.integrations.agent_identity._iam_connector_credentials_provider import Client +from google.adk.sessions.session import Session +from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsMetadata +from google.cloud.iamconnectorcredentials_v1alpha import RetrieveCredentialsResponse +from google.longrunning.operations_pb2 import Operation + + +@pytest.fixture +def mock_client(): + return Mock(spec=Client) + + +@pytest.fixture +def provider(mock_client): + return _IamConnectorCredentialsProvider(client=mock_client) + + +@pytest.fixture +def auth_scheme(): + scheme = GcpAuthProviderScheme( + name="projects/test-project/locations/global/connectors/test-connector", + scopes=["test-scope"], + continue_uri="https://example.com/continue", + ) + return scheme + + +@pytest.fixture +def mock_operation(mock_client): + op = Operation(done=True) + + class DummyCall: + + def __init__(self, operation): + self.operation = operation + + mock_client.retrieve_credentials.return_value = DummyCall(op) + return op + + +@pytest.fixture +def context(): + context = Mock(spec=CallbackContext) + context.user_id = "user" + context.function_call_id = "call_123" + session = Mock(spec=Session) + session.events = [] + context.session = session + + return context + + + +@patch.dict(_iam_connector_credentials_provider.os.environ, clear=True) +@patch.object(_iam_connector_credentials_provider, "Client") +def test_get_client_uses_rest_transport(mock_client_class): + provider = _iam_connector_credentials_provider._IamConnectorCredentialsProvider() + provider._get_client() + + mock_client_class.assert_called_once() + _, kwargs = mock_client_class.call_args + assert kwargs.get("transport") == "rest" + + +@patch.dict( + _iam_connector_credentials_provider.os.environ, + {"IAM_CONNECTOR_CREDENTIALS_TARGET_HOST": "some-host"}, +) +@patch.object(_iam_connector_credentials_provider, "Client") +@patch.object(_iam_connector_credentials_provider, "ClientOptions") +def test_get_client_with_env_var(mock_client_options_class, mock_client_class): + provider = _iam_connector_credentials_provider._IamConnectorCredentialsProvider() + client = provider._get_client() + + assert client == mock_client_class.return_value + mock_client_options_class.assert_called_once_with(api_endpoint="some-host") + mock_client_class.assert_called_once_with( + client_options=mock_client_options_class.return_value, transport="rest" + ) + + +# ============================================================================== +# Non-interactive auth flows (API key and 2-legged OAuth) +# ============================================================================== + + +async def test_get_auth_credential_raises_error_if_context_is_missing( + provider, auth_scheme +): + """Test get_auth_credential raises ValueError if context is missing.""" + with pytest.raises( + ValueError, + match="GcpAuthProvider requires a context with a valid user_id", + ): + await provider.get_auth_credential(auth_scheme, context=None) + + +async def test_get_auth_credential_raises_error_if_user_id_is_missing( + provider, auth_scheme +): + """Test get_auth_credential raises ValueError if user_id is missing.""" + context = Mock(spec=CallbackContext) + context.user_id = None + with pytest.raises( + ValueError, + match="GcpAuthProvider requires a context with a valid user_id", + ): + await provider.get_auth_credential(auth_scheme, context=context) + + +async def test_get_auth_credential_returns_credential_if_available_immediately( + mock_client, + mock_operation, + auth_scheme, + context, + provider, +): + """Test get_auth_credential returns credential if available immediately.""" + mock_credential = RetrieveCredentialsResponse( + header="Authorization: Bearer", token="test-token" + ) + mock_operation.response.value = RetrieveCredentialsResponse.serialize( + mock_credential + ) + + auth_credential = await provider.get_auth_credential(auth_scheme, context) + + assert auth_credential.auth_type == AuthCredentialTypes.HTTP + assert auth_credential.http.scheme == "Bearer" + assert auth_credential.http.credentials.token == "test-token" + mock_client.retrieve_credentials.assert_called_once() + + +async def test_get_auth_credential_raises_error_if_upstream_returns_empty_header( + mock_operation, + auth_scheme, + context, + provider, +): + """Test get_auth_credential raises RuntimeError for empty header.""" + mock_credential = RetrieveCredentialsResponse(header="", token="test-token") + mock_operation.response.value = RetrieveCredentialsResponse.serialize( + mock_credential + ) + + with pytest.raises( + ValueError, + match=( + "Received either empty header or token from IAM Connector" + " Credentials service." + ), + ): + await provider.get_auth_credential(auth_scheme, context) + + +async def test_get_auth_credential_raises_error_if_upstream_returns_empty_token( + mock_operation, + auth_scheme, + context, + provider, +): + """Test get_auth_credential raises RuntimeError for empty token.""" + mock_credential = RetrieveCredentialsResponse( + header="Authorization: Bearer", token="" + ) + mock_operation.response.value = RetrieveCredentialsResponse.serialize( + mock_credential + ) + + with pytest.raises( + ValueError, + match=( + "Received either empty header or token from IAM Connector" + " Credentials service." + ), + ): + await provider.get_auth_credential(auth_scheme, context) + + +async def test_get_auth_credential_returns_credential_if_upstream_returns_custom_header( + mock_operation, + auth_scheme, + context, + provider, +): + """Test get_auth_credential returns valid credential for custom header and sets X-GOOG-API-KEY header.""" + mock_credential = RetrieveCredentialsResponse( + header="some-x-api-key", token="test-token" + ) + mock_operation.response.value = RetrieveCredentialsResponse.serialize( + mock_credential + ) + + auth_credential = await provider.get_auth_credential(auth_scheme, context) + + assert auth_credential.auth_type == AuthCredentialTypes.HTTP + assert not auth_credential.http.scheme + assert auth_credential.http.credentials.token is None + assert auth_credential.http.additional_headers == { + "some-x-api-key": "test-token", + "X-GOOG-API-KEY": "test-token", + } + + +async def test_get_auth_credential_raises_error_if_upstream_operation_errors( + mock_operation, auth_scheme, context, provider +): + """Test get_auth_credential raises RuntimeError for failed operations.""" + mock_operation.error.message = "OAuth server error" + mock_operation.done = False + + with pytest.raises( + RuntimeError, match="Operation failed: OAuth server error" + ): + await provider.get_auth_credential(auth_scheme, context) + + +async def test_get_auth_credential_raises_error_if_upstream_call_fails( + mock_client, auth_scheme, context, provider +): + """Test get_auth_credential raises RuntimeError for failed calls.""" + mock_client.retrieve_credentials.side_effect = Exception( + "API Quota Exhausted" + ) + + with pytest.raises( + RuntimeError, + match="Failed to retrieve credential for user 'user' on connector", + ) as exc_info: + await provider.get_auth_credential(auth_scheme, context) + + # Assert that the original Exception is the chained cause! + assert str(exc_info.value.__cause__) == "API Quota Exhausted" + + +@patch.object(_iam_connector_credentials_provider.time, "time") +async def test_get_auth_credential_raises_error_if_polling_times_out( + mock_time, + mock_operation, + auth_scheme, + context, + provider, +): + """Test get_auth_credential raises RuntimeError if polling times out.""" + + # Force the operation into the polling loop state + meta_pb = RetrieveCredentialsMetadata.pb()() + meta_pb.consent_pending.SetInParent() + meta = RetrieveCredentialsMetadata.deserialize(meta_pb.SerializeToString()) + mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta) + + # First call sets start_time=0.0, second call checks time > timeout + # (20.0 > 10.0) + mock_time.side_effect = [0.0, 20.0] + + mock_metadata = Mock(spec=RetrieveCredentialsMetadata) + mock_metadata.consent_pending = True + mock_metadata.uri_consent_required = False + mock_operation.done = True + mock_operation.ClearField("error") + mock_client = Mock(spec=Client) + mock_client.retrieve_credentials.side_effect = Exception( + "Timeout waiting for credentials." + ) + provider._client = mock_client + + with pytest.raises( + RuntimeError, + match="Failed to retrieve credential for user 'user' on connector", + ) as exc_info: + await provider.get_auth_credential(auth_scheme, context) + + assert "Timeout waiting for credentials." in str(exc_info.value.__cause__) + + +# ============================================================================== +# Interactive Auth Flows (3-legged OAuth for User Consents) +# ============================================================================== + + +async def test_get_auth_credential_initiates_user_consent( + mock_operation, auth_scheme, context, provider +): + # Explicitly set the mock behavior for this test + expected_uri = "https://example.com/auth" + expected_nonce = "sample-nonce-123" + meta = RetrieveCredentialsMetadata({ + "uri_consent_required": { + "authorization_uri": expected_uri, + "consent_nonce": expected_nonce, + } + }) + mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta) + mock_operation.done = False + # Assert that there is no prior user consent completion event + assert not context.session.events + + credential = await provider.get_auth_credential(auth_scheme, context) + + assert credential is not None + assert credential.auth_type == AuthCredentialTypes.OAUTH2 + assert credential.oauth2.auth_uri == expected_uri + assert credential.oauth2.nonce == expected_nonce + + +async def test_get_auth_credential_returns_fresh_auth_uri_for_repeated_requests( + mock_client, mock_operation, auth_scheme, context, provider +): + """Test that repeated calls fetch fresh auth URIs if consent is still pending.""" + # Arrange: Explicit initial URI + initial_uri = "https://example.com/auth" + initial_nonce = "initial-nonce-123" + meta1 = RetrieveCredentialsMetadata({ + "uri_consent_required": { + "authorization_uri": initial_uri, + "consent_nonce": initial_nonce, + } + }) + mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta1) + mock_operation.done = False + + credential1 = await provider.get_auth_credential(auth_scheme, context) + assert credential1.oauth2.auth_uri == initial_uri + assert credential1.oauth2.nonce == initial_nonce + + # Arrange: Explicit new URI for the second call + fresh_auth_uri = "https://example.com/auth_new" + fresh_nonce = "fresh-nonce-456" + meta2 = RetrieveCredentialsMetadata({ + "uri_consent_required": { + "authorization_uri": fresh_auth_uri, + "consent_nonce": fresh_nonce, + } + }) + mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta2) + + credential2 = await provider.get_auth_credential(auth_scheme, context) + + assert mock_client.retrieve_credentials.call_count == 2 + assert credential2.oauth2.auth_uri == fresh_auth_uri + assert credential2.oauth2.nonce == fresh_nonce + + +async def test_get_auth_credential_returns_token_if_consent_was_completed( + mock_operation, auth_scheme, context, provider +): + # Setup mock credential for successful credential retrieval + mock_credential = RetrieveCredentialsResponse( + header="Authorization: Bearer", token="test-token" + ) + mock_operation.response.value = RetrieveCredentialsResponse.serialize( + mock_credential + ) + + # Create mock events + # 1. FunctionCall event for adk_request_credential + function_call = Mock() + function_call.id = "auth-req-1" + function_call.name = REQUEST_EUC_FUNCTION_CALL_NAME + function_call.args = AuthToolArguments( + function_call_id="call-123", auth_config=Mock(spec=AuthConfig, auth_scheme=auth_scheme) + ).model_dump(by_alias=True, exclude_none=True) + + event1 = Mock() + event1.get_function_calls.return_value = [function_call] + event1.get_function_responses.return_value = [] + + # 2. FunctionResponse event for adk_request_credential + function_response = Mock() + function_response.id = "auth-req-1" + function_response.name = REQUEST_EUC_FUNCTION_CALL_NAME + + event2 = Mock() + event2.get_function_calls.return_value = [] + event2.get_function_responses.return_value = [function_response] + + # Setup tool context and event history (order of events matters) + context.session.events = [event1, event2] + context.function_call_id = "call-123" + + # Also set uri_consent_required to True-ish so it enters the check block + meta = RetrieveCredentialsMetadata( + uri_consent_required=RetrieveCredentialsMetadata.UriConsentRequired() + ) + mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta) + + # Execute + auth_credential = await provider.get_auth_credential(auth_scheme, context) + + # Verify + assert auth_credential is not None + assert auth_credential.auth_type == AuthCredentialTypes.HTTP + assert auth_credential.http.scheme == "Bearer" + assert auth_credential.http.credentials.token == "test-token" + + +async def test_get_auth_credential_raises_error_if_consent_canceled( + mock_operation, auth_scheme, context, provider +): + function_call = Mock() + function_call.id = "auth-req-1" + function_call.name = REQUEST_EUC_FUNCTION_CALL_NAME + function_call.args = AuthToolArguments( + function_call_id="call-123", auth_config=Mock(spec=AuthConfig, auth_scheme=auth_scheme) + ).model_dump(by_alias=True, exclude_none=True) + + event1 = Mock() + event1.get_function_calls.return_value = [function_call] + event1.get_function_responses.return_value = [] + + function_response = Mock() + function_response.id = "auth-req-1" + function_response.name = REQUEST_EUC_FUNCTION_CALL_NAME + + event2 = Mock() + event2.get_function_calls.return_value = [] + event2.get_function_responses.return_value = [function_response] + + context.session.events = [event1, event2] + context.function_call_id = "call-123" + + meta = RetrieveCredentialsMetadata({ + "uri_consent_required": { + "authorization_uri": "https://example.com/auth", + "consent_nonce": "sample-nonce", + } + }) + mock_operation.metadata.value = RetrieveCredentialsMetadata.serialize(meta) + mock_operation.done = False + + with pytest.raises( + RuntimeError, match="Failed to retrieve consent based credential." + ): + await provider.get_auth_credential(auth_scheme, context)