From 989721746aba65e90f644e51606375699175f709 Mon Sep 17 00:00:00 2001 From: Aarav Mittal <137450929+a2105z@users.noreply.github.com> Date: Mon, 17 Aug 2026 19:00:10 -0700 Subject: [PATCH] fix: use OAuth2 client-credentials scheme for OpenAPI SA helpers Merge https://github.com/google/adk-python/pull/6660 Change service-account OpenAPI helpers to return an OAuth2 client-credentials scheme so CredentialManager can perform token exchange. Also, bypass the credential service caching for all SERVICE_ACCOUNT credentials. This ensures we don't cache exchanged tokens that cannot be refreshed, but means token exchange will run on each tool execution if the manager/exchanger is not reused. Fixes #6656 PiperOrigin-RevId: 966305100 --- scripts/compliance_checks.py | 1 - src/google/adk/auth/credential_manager.py | 14 +- .../tools/openapi_tool/auth/auth_helpers.py | 41 ++- .../service_account_exchanger.py | 79 +++++- .../unittests/auth/test_credential_manager.py | 64 +++++ .../test_service_account_exchanger.py | 247 ++++++++++++++++++ .../openapi_tool/auth/test_auth_helper.py | 56 +++- 7 files changed, 487 insertions(+), 15 deletions(-) diff --git a/scripts/compliance_checks.py b/scripts/compliance_checks.py index 285efeeb..6d5a7030 100755 --- a/scripts/compliance_checks.py +++ b/scripts/compliance_checks.py @@ -40,7 +40,6 @@ _EXCLUDED_FROM_MTLS = { 'src/google/adk/tools/_google_credentials.py', 'src/google/adk/tools/apihub_tool/clients/apihub_client.py', 'src/google/adk/tools/google_api_tool/google_api_toolset.py', - 'src/google/adk/tools/openapi_tool/auth/auth_helpers.py', 'tests/unittests/auth/test_credential_manager.py', 'tests/unittests/flows/llm_flows/test_functions_request_euc.py', 'tests/unittests/tools/application_integration_tool/test_application_integration_toolset.py', diff --git a/src/google/adk/auth/credential_manager.py b/src/google/adk/auth/credential_manager.py index d693f672..3135d745 100644 --- a/src/google/adk/auth/credential_manager.py +++ b/src/google/adk/auth/credential_manager.py @@ -238,7 +238,12 @@ class CredentialManager: return raw_auth_credential.model_copy(deep=True) # Step 3: Try to load existing processed credential - credential = await self._load_existing_credential(context) + credential = None + if not ( + raw_auth_credential + and raw_auth_credential.auth_type == AuthCredentialTypes.SERVICE_ACCOUNT + ): + credential = await self._load_existing_credential(context) # Step 4: If no existing credential, load from auth response # TODO instead of load from auth response, we can store auth response in @@ -269,7 +274,12 @@ class CredentialManager: # Step 8: Save credential if it was modified if was_from_auth_response or was_exchanged or was_refreshed: - await self._save_credential(context, credential) + if not ( + raw_auth_credential + and raw_auth_credential.auth_type + == AuthCredentialTypes.SERVICE_ACCOUNT + ): + await self._save_credential(context, credential) return credential diff --git a/src/google/adk/tools/openapi_tool/auth/auth_helpers.py b/src/google/adk/tools/openapi_tool/auth/auth_helpers.py index 2c8ae5bb..78175e6b 100644 --- a/src/google/adk/tools/openapi_tool/auth/auth_helpers.py +++ b/src/google/adk/tools/openapi_tool/auth/auth_helpers.py @@ -26,6 +26,8 @@ from fastapi.openapi.models import APIKeyIn from fastapi.openapi.models import HTTPBase from fastapi.openapi.models import HTTPBearer from fastapi.openapi.models import OAuth2 +from fastapi.openapi.models import OAuthFlowClientCredentials +from fastapi.openapi.models import OAuthFlows from fastapi.openapi.models import OpenIdConnect from fastapi.openapi.models import Schema import httpx @@ -152,13 +154,40 @@ def token_to_scheme_credential( raise ValueError(f"Invalid security scheme type: {type}") +def _service_account_auth_scheme() -> OAuth2: + """Auth scheme for Google Service Account credentials. + + CredentialManager only auto-loads raw non-interactive credentials when the + scheme is an OAuth2/OIDC client-credentials flow. An HTTPBearer scheme makes + ``_is_client_credentials_flow`` return False, so ``get_auth_credential`` + returns None and the tool falls back to ``adk_request_credential`` instead of + exchanging the service account for a token. + + The token URL is unused by ServiceAccountCredentialExchanger (ADC / JWT + assertion), but is required by the OAuth2 client-credentials model. + """ + return OAuth2( + flows=OAuthFlows( + clientCredentials=OAuthFlowClientCredentials( + # Placeholder only; SA exchange does not call this endpoint. + # Use the mTLS host form for compliance with Google API endpoint + # requirements. + tokenUrl="https://oauth2.mtls.googleapis.com/token", + scopes={}, + ) + ) + ) + + def service_account_dict_to_scheme_credential( config: Dict[str, Any], scopes: List[str], ) -> Tuple[AuthScheme, AuthCredential]: """Creates AuthScheme and AuthCredential for Google Service Account. - Returns a bearer token scheme, and a service account credential. + Returns an OAuth2 client-credentials scheme (so CredentialManager can + exchange the service account) and a service account credential. After + exchange the credential is an HTTP bearer token. Args: config: A ServiceAccount object containing the Google Service Account @@ -168,7 +197,6 @@ def service_account_dict_to_scheme_credential( Returns: Tuple: (AuthScheme, AuthCredential) """ - auth_scheme = HTTPBearer(bearerFormat="JWT") service_account = ServiceAccount( service_account_credential=ServiceAccountCredential.model_construct( **config @@ -179,7 +207,7 @@ def service_account_dict_to_scheme_credential( auth_type=AuthCredentialTypes.SERVICE_ACCOUNT, service_account=service_account, ) - return auth_scheme, auth_credential + return _service_account_auth_scheme(), auth_credential def service_account_scheme_credential( @@ -187,7 +215,9 @@ def service_account_scheme_credential( ) -> Tuple[AuthScheme, AuthCredential]: """Creates AuthScheme and AuthCredential for Google Service Account. - Returns a bearer token scheme, and a service account credential. + Returns an OAuth2 client-credentials scheme (so CredentialManager can + exchange the service account) and a service account credential. After + exchange the credential is an HTTP bearer token. Args: config: A ServiceAccount object containing the Google Service Account @@ -196,11 +226,10 @@ def service_account_scheme_credential( Returns: Tuple: (AuthScheme, AuthCredential) """ - auth_scheme = HTTPBearer(bearerFormat="JWT") auth_credential = AuthCredential( auth_type=AuthCredentialTypes.SERVICE_ACCOUNT, service_account=config ) - return auth_scheme, auth_credential + return _service_account_auth_scheme(), auth_credential def openid_dict_to_scheme_credential( diff --git a/src/google/adk/tools/openapi_tool/auth/credential_exchangers/service_account_exchanger.py b/src/google/adk/tools/openapi_tool/auth/credential_exchangers/service_account_exchanger.py index 2b79edf9..f753cc40 100644 --- a/src/google/adk/tools/openapi_tool/auth/credential_exchangers/service_account_exchanger.py +++ b/src/google/adk/tools/openapi_tool/auth/credential_exchangers/service_account_exchanger.py @@ -16,10 +16,14 @@ from __future__ import annotations +import calendar +import time +from typing import Any from typing import Optional import google.auth from google.auth import exceptions as google_auth_exceptions +from google.auth import jwt from google.auth.transport.requests import Request from google.oauth2 import service_account import google.oauth2.credentials @@ -33,6 +37,38 @@ from .....auth.auth_schemes import AuthScheme from .base_credential_exchanger import AuthCredentialMissingError from .base_credential_exchanger import BaseAuthCredentialExchanger +_access_token_cache: dict[tuple[Any, ...], tuple[AuthCredential, float]] = {} +_id_token_cache: dict[tuple[Any, ...], tuple[AuthCredential, float]] = {} + + +def _get_cache_key(sa_config: ServiceAccount) -> tuple[Any, ...]: + scopes_tuple = tuple(sa_config.scopes) if sa_config.scopes else () + if sa_config.use_default_credential: + return ( + True, + scopes_tuple, + sa_config.use_id_token, + sa_config.audience, + ) + else: + cred = sa_config.service_account_credential + cred_id = cred.private_key_id if cred else None + client_email = cred.client_email if cred else None + return ( + False, + cred_id, + client_email, + scopes_tuple, + sa_config.use_id_token, + sa_config.audience, + ) + + +def _reset_cache(): + global _access_token_cache, _id_token_cache + _access_token_cache.clear() + _id_token_cache.clear() + class ServiceAccountCredentialExchanger(BaseAuthCredentialExchanger): """Fetches credentials for Google Service Account. @@ -95,6 +131,13 @@ class ServiceAccountCredentialExchanger(BaseAuthCredentialExchanger): Raises: AuthCredentialMissingError: If token exchange fails. """ + cache_key = _get_cache_key(sa_config) + cached_val = _id_token_cache.get(cache_key) + if cached_val: + token, expires_at = cached_val + if time.time() < expires_at - 300: + return token + # audience and credential presence are validated by the ServiceAccount # model_validator at construction time. try: @@ -103,6 +146,11 @@ class ServiceAccountCredentialExchanger(BaseAuthCredentialExchanger): request = Request() token = oauth2_id_token.fetch_id_token(request, sa_config.audience) + try: + decoded = jwt.decode(token, verify=False) + expires_at = decoded.get("exp") or int(time.time() + 3600) + except Exception: # pylint: disable=broad-except + expires_at = int(time.time() + 3600) else: # Guaranteed non-None by ServiceAccount model_validator. assert sa_config.service_account_credential is not None @@ -114,14 +162,24 @@ class ServiceAccountCredentialExchanger(BaseAuthCredentialExchanger): ) credentials.refresh(Request()) token = credentials.token + try: + expires_at = ( + calendar.timegm(credentials.expiry.utctimetuple()) + if credentials.expiry + else int(time.time() + 3600) + ) + except (AttributeError, TypeError, ValueError): + expires_at = int(time.time() + 3600) - return AuthCredential( + res = AuthCredential( auth_type=AuthCredentialTypes.HTTP, http=HttpAuth( scheme="bearer", credentials=HttpCredentials(token=token), ), ) + _id_token_cache[cache_key] = (res, expires_at) + return res # ValueError is raised by google-auth when service account JSON is # missing required fields (e.g. client_email, private_key), or when @@ -146,6 +204,13 @@ class ServiceAccountCredentialExchanger(BaseAuthCredentialExchanger): AuthCredentialMissingError: If scopes are missing for explicit credentials or token exchange fails. """ + cache_key = _get_cache_key(sa_config) + cached_val = _access_token_cache.get(cache_key) + if cached_val: + token, expires_at = cached_val + if time.time() < expires_at - 300: + return token + if not sa_config.use_default_credential and not sa_config.scopes: raise AuthCredentialMissingError( "scopes are required when using explicit service account credentials" @@ -173,8 +238,16 @@ class ServiceAccountCredentialExchanger(BaseAuthCredentialExchanger): quota_project_id = None credentials.refresh(Request()) + try: + expires_at = ( + calendar.timegm(credentials.expiry.utctimetuple()) + if credentials.expiry + else int(time.time() + 3600) + ) + except (AttributeError, TypeError, ValueError): + expires_at = int(time.time() + 3600) - return AuthCredential( + res = AuthCredential( auth_type=AuthCredentialTypes.HTTP, http=HttpAuth( scheme="bearer", @@ -186,6 +259,8 @@ class ServiceAccountCredentialExchanger(BaseAuthCredentialExchanger): else None, ), ) + _access_token_cache[cache_key] = (res, expires_at) + return res # ValueError is raised by google-auth when service account JSON is # missing required fields (e.g. client_email, private_key). diff --git a/tests/unittests/auth/test_credential_manager.py b/tests/unittests/auth/test_credential_manager.py index 0ba94908..18dc8bb7 100644 --- a/tests/unittests/auth/test_credential_manager.py +++ b/tests/unittests/auth/test_credential_manager.py @@ -329,6 +329,70 @@ class TestCredentialManager: assert result is None + @pytest.mark.asyncio + async def test_get_auth_credential_service_account_skips_cache( + self, mocker, service_account_credential + ): + """Test that Service Account credentials bypass the load/save cache.""" + from fastapi.openapi.models import OAuth2 + from fastapi.openapi.models import OAuthFlowClientCredentials + from fastapi.openapi.models import OAuthFlows + from google.adk.auth.auth_credential import HttpAuth + from google.adk.auth.auth_credential import HttpCredentials + + auth_scheme = OAuth2( + flows=OAuthFlows( + clientCredentials=OAuthFlowClientCredentials( + tokenUrl="https://example.com/token", + scopes={}, + ) + ) + ) + + auth_config = AuthConfig( + auth_scheme=auth_scheme, + raw_auth_credential=service_account_credential, + ) + + exchanged_credential = AuthCredential( + auth_type=AuthCredentialTypes.HTTP, + http=HttpAuth( + scheme="bearer", + credentials=HttpCredentials(token="sa-access-token"), + ), + ) + + tool_context = mocker.Mock(spec=CallbackContext) + + manager = CredentialManager(auth_config) + + # Mock the private methods + manager._validate_credential = mocker.AsyncMock() + manager._is_credential_ready = mocker.Mock(return_value=False) + manager._load_existing_credential = mocker.AsyncMock() + manager._load_from_auth_response = mocker.AsyncMock(return_value=None) + manager._exchange_credential = mocker.AsyncMock( + return_value=(exchanged_credential, True) + ) + manager._refresh_credential = mocker.AsyncMock( + return_value=(exchanged_credential, False) + ) + manager._save_credential = mocker.AsyncMock() + manager._is_client_credentials_flow = mocker.Mock(return_value=True) + + result = await manager.get_auth_credential(tool_context) + + # Verify load and save were NOT called + manager._load_existing_credential.assert_not_called() + manager._save_credential.assert_not_called() + + # Verify exchange WAS called + manager._exchange_credential.assert_called_once() + called_arg = manager._exchange_credential.call_args[0][0] + assert called_arg.auth_type == AuthCredentialTypes.SERVICE_ACCOUNT + + assert result == exchanged_credential + @pytest.mark.asyncio async def test_load_existing_credential_already_exchanged(self): """Test _load_existing_credential ignores shared config cache.""" diff --git a/tests/unittests/tools/openapi_tool/auth/credential_exchangers/test_service_account_exchanger.py b/tests/unittests/tools/openapi_tool/auth/credential_exchangers/test_service_account_exchanger.py index fb35daf6..26365264 100644 --- a/tests/unittests/tools/openapi_tool/auth/credential_exchangers/test_service_account_exchanger.py +++ b/tests/unittests/tools/openapi_tool/auth/credential_exchangers/test_service_account_exchanger.py @@ -14,6 +14,9 @@ """Unit tests for the service account credential exchanger.""" +import calendar +import datetime +import time from unittest.mock import MagicMock from google.adk.auth.auth_credential import AuthCredential @@ -23,6 +26,7 @@ from google.adk.auth.auth_credential import ServiceAccountCredential from google.adk.auth.auth_schemes import AuthScheme from google.adk.auth.auth_schemes import AuthSchemeType from google.adk.tools.openapi_tool.auth.credential_exchangers.base_credential_exchanger import AuthCredentialMissingError +from google.adk.tools.openapi_tool.auth.credential_exchangers.service_account_exchanger import _reset_cache from google.adk.tools.openapi_tool.auth.credential_exchangers.service_account_exchanger import ServiceAccountCredentialExchanger import google.auth from google.auth import exceptions as google_auth_exceptions @@ -43,6 +47,11 @@ _ID_TOKEN_MONKEYPATCH_TARGET = ( _FETCH_ID_TOKEN_MONKEYPATCH_TARGET = "google.oauth2.id_token.fetch_id_token" +@pytest.fixture(autouse=True) +def reset_exchanger_cache(): + _reset_cache() + + @pytest.fixture def service_account_exchanger(): return ServiceAccountCredentialExchanger() @@ -391,3 +400,241 @@ def test_model_validator_allows_adc_without_explicit_credential(): ) assert sa.service_account_credential is None assert sa.use_default_credential is True + + +def test_exchange_access_token_caching( + service_account_exchanger, auth_scheme, sa_credential, monkeypatch +): + mock_credentials = MagicMock() + mock_credentials.token = "mock_access_token" + mock_credentials.expiry = datetime.datetime(2026, 8, 10, 22, 0, 0) + mock_credentials.quota_project_id = None + + mock_from_sa_info = MagicMock(return_value=mock_credentials) + monkeypatch.setattr(_ACCESS_TOKEN_MONKEYPATCH_TARGET, mock_from_sa_info) + + expiry_timestamp = calendar.timegm(mock_credentials.expiry.utctimetuple()) + current_time = expiry_timestamp - 600 + monkeypatch.setattr(time, "time", lambda: current_time) + + auth_credential = AuthCredential( + auth_type=AuthCredentialTypes.SERVICE_ACCOUNT, + service_account=ServiceAccount( + service_account_credential=sa_credential, + scopes=_DEFAULT_SCOPES, + ), + ) + + # First call - should exchange + result1 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result1.http.credentials.token == "mock_access_token" + assert mock_from_sa_info.call_count == 1 + + # Second call - should return cached + result2 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result2.http.credentials.token == "mock_access_token" + assert mock_from_sa_info.call_count == 1 + + # Third call - time moves forward, close to expiry + current_time = expiry_timestamp - 200 # 200s < 300s, so expired + monkeypatch.setattr(time, "time", lambda: current_time) + + mock_credentials2 = MagicMock() + mock_credentials2.token = "new_mock_access_token" + mock_credentials2.expiry = datetime.datetime(2026, 8, 10, 23, 0, 0) + mock_credentials2.quota_project_id = None + mock_from_sa_info.return_value = mock_credentials2 + + result3 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result3.http.credentials.token == "new_mock_access_token" + assert mock_from_sa_info.call_count == 2 + + # Fourth call - config changes + auth_credential_new_scopes = AuthCredential( + auth_type=AuthCredentialTypes.SERVICE_ACCOUNT, + service_account=ServiceAccount( + service_account_credential=sa_credential, + scopes=["another-scope"], + ), + ) + mock_credentials3 = MagicMock() + mock_credentials3.token = "another_scope_token" + mock_credentials3.expiry = datetime.datetime(2026, 8, 11, 0, 0, 0) + mock_credentials3.quota_project_id = None + mock_from_sa_info.return_value = mock_credentials3 + + result4 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential_new_scopes + ) + assert result4.http.credentials.token == "another_scope_token" + assert mock_from_sa_info.call_count == 3 + + +def test_exchange_id_token_caching_explicit( + service_account_exchanger, auth_scheme, sa_credential, monkeypatch +): + mock_id_credentials = MagicMock() + mock_id_credentials.token = "mock_id_token" + mock_id_credentials.expiry = datetime.datetime(2026, 8, 10, 22, 0, 0) + mock_from_sa_info = MagicMock(return_value=mock_id_credentials) + monkeypatch.setattr(_ID_TOKEN_MONKEYPATCH_TARGET, mock_from_sa_info) + + expiry_timestamp = calendar.timegm(mock_id_credentials.expiry.utctimetuple()) + current_time = expiry_timestamp - 600 + monkeypatch.setattr(time, "time", lambda: current_time) + + auth_credential = AuthCredential( + auth_type=AuthCredentialTypes.SERVICE_ACCOUNT, + service_account=ServiceAccount( + service_account_credential=sa_credential, + scopes=_DEFAULT_SCOPES, + use_id_token=True, + audience="https://my-service.run.app", + ), + ) + + # First call + result1 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result1.http.credentials.token == "mock_id_token" + assert mock_from_sa_info.call_count == 1 + + # Second call - cached + result2 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result2.http.credentials.token == "mock_id_token" + assert mock_from_sa_info.call_count == 1 + + # Third call - expired + current_time = expiry_timestamp - 200 + monkeypatch.setattr(time, "time", lambda: current_time) + + mock_id_credentials2 = MagicMock() + mock_id_credentials2.token = "new_mock_id_token" + mock_id_credentials2.expiry = datetime.datetime(2026, 8, 10, 23, 0, 0) + mock_from_sa_info.return_value = mock_id_credentials2 + + result3 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result3.http.credentials.token == "new_mock_id_token" + assert mock_from_sa_info.call_count == 2 + + +def test_exchange_id_token_caching_adc( + service_account_exchanger, auth_scheme, monkeypatch +): + mock_fetch_id_token = MagicMock(return_value="mock_adc_id_token") + monkeypatch.setattr(_FETCH_ID_TOKEN_MONKEYPATCH_TARGET, mock_fetch_id_token) + + expiry_timestamp = 1773268800 # 2026-08-10 22:00:00 UTC + mock_jwt_decode = MagicMock(return_value={"exp": expiry_timestamp}) + monkeypatch.setattr( + "google.adk.tools.openapi_tool.auth.credential_exchangers.service_account_exchanger.jwt.decode", + mock_jwt_decode, + ) + + current_time = expiry_timestamp - 600 + monkeypatch.setattr(time, "time", lambda: current_time) + + auth_credential = AuthCredential( + auth_type=AuthCredentialTypes.SERVICE_ACCOUNT, + service_account=ServiceAccount( + use_default_credential=True, + scopes=_DEFAULT_SCOPES, + use_id_token=True, + audience="https://my-service.run.app", + ), + ) + + # First call + result1 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result1.http.credentials.token == "mock_adc_id_token" + assert mock_fetch_id_token.call_count == 1 + mock_jwt_decode.assert_called_once_with("mock_adc_id_token", verify=False) + + # Second call - cached + result2 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result2.http.credentials.token == "mock_adc_id_token" + assert mock_fetch_id_token.call_count == 1 + + # Third call - expired + current_time = expiry_timestamp - 200 + monkeypatch.setattr(time, "time", lambda: current_time) + + mock_fetch_id_token.return_value = "new_mock_adc_id_token" + mock_jwt_decode.return_value = {"exp": expiry_timestamp + 3600} + + result3 = service_account_exchanger.exchange_credential( + auth_scheme, auth_credential + ) + assert result3.http.credentials.token == "new_mock_adc_id_token" + assert mock_fetch_id_token.call_count == 2 + + +def test_exchange_access_token_caching_different_client_emails( + service_account_exchanger, auth_scheme, sa_credential, monkeypatch +): + mock_credentials = MagicMock() + mock_credentials.token = "mock_access_token_1" + mock_credentials.expiry = datetime.datetime(2026, 8, 10, 22, 0, 0) + mock_credentials.quota_project_id = None + + mock_from_sa_info = MagicMock(return_value=mock_credentials) + monkeypatch.setattr(_ACCESS_TOKEN_MONKEYPATCH_TARGET, mock_from_sa_info) + + expiry_timestamp = calendar.timegm(mock_credentials.expiry.utctimetuple()) + current_time = expiry_timestamp - 600 + monkeypatch.setattr(time, "time", lambda: current_time) + + sa_cred_1 = sa_credential.model_copy( + update={"private_key_id": None, "client_email": "sa1@example.com"} + ) + sa_cred_2 = sa_credential.model_copy( + update={"private_key_id": None, "client_email": "sa2@example.com"} + ) + + auth_cred_1 = AuthCredential( + auth_type=AuthCredentialTypes.SERVICE_ACCOUNT, + service_account=ServiceAccount( + service_account_credential=sa_cred_1, + scopes=_DEFAULT_SCOPES, + ), + ) + auth_cred_2 = AuthCredential( + auth_type=AuthCredentialTypes.SERVICE_ACCOUNT, + service_account=ServiceAccount( + service_account_credential=sa_cred_2, + scopes=_DEFAULT_SCOPES, + ), + ) + + result1 = service_account_exchanger.exchange_credential( + auth_scheme, auth_cred_1 + ) + assert result1.http.credentials.token == "mock_access_token_1" + assert mock_from_sa_info.call_count == 1 + + mock_credentials_2 = MagicMock() + mock_credentials_2.token = "mock_access_token_2" + mock_credentials_2.expiry = datetime.datetime(2026, 8, 10, 22, 0, 0) + mock_credentials_2.quota_project_id = None + mock_from_sa_info.return_value = mock_credentials_2 + + result2 = service_account_exchanger.exchange_credential( + auth_scheme, auth_cred_2 + ) + assert result2.http.credentials.token == "mock_access_token_2" + assert mock_from_sa_info.call_count == 2 diff --git a/tests/unittests/tools/openapi_tool/auth/test_auth_helper.py b/tests/unittests/tools/openapi_tool/auth/test_auth_helper.py index 3f5e8f07..8a569f19 100644 --- a/tests/unittests/tools/openapi_tool/auth/test_auth_helper.py +++ b/tests/unittests/tools/openapi_tool/auth/test_auth_helper.py @@ -12,6 +12,8 @@ # 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 from fastapi.openapi.models import APIKey @@ -28,6 +30,8 @@ from google.adk.auth.auth_credential import ServiceAccount from google.adk.auth.auth_credential import ServiceAccountCredential from google.adk.auth.auth_schemes import AuthSchemeType from google.adk.auth.auth_schemes import OpenIdConnectWithConfig +from google.adk.auth.auth_tool import AuthConfig +from google.adk.auth.credential_manager import CredentialManager from google.adk.tools.openapi_tool.auth.auth_helpers import credential_to_param from google.adk.tools.openapi_tool.auth.auth_helpers import dict_to_auth_scheme from google.adk.tools.openapi_tool.auth.auth_helpers import INTERNAL_AUTH_PREFIX @@ -133,8 +137,10 @@ def test_service_account_dict_to_scheme_credential(): scheme, credential = service_account_dict_to_scheme_credential(config, scopes) - assert isinstance(scheme, HTTPBearer) - assert scheme.bearerFormat == "JWT" + assert isinstance(scheme, OAuth2) + assert scheme.flows is not None + assert scheme.flows.clientCredentials is not None + assert scheme.flows.clientCredentials.tokenUrl assert credential.auth_type == AuthCredentialTypes.SERVICE_ACCOUNT assert credential.service_account.scopes == scopes assert ( @@ -163,12 +169,54 @@ def test_service_account_scheme_credential(): scheme, credential = service_account_scheme_credential(config) - assert isinstance(scheme, HTTPBearer) - assert scheme.bearerFormat == "JWT" + assert isinstance(scheme, OAuth2) + assert scheme.flows is not None + assert scheme.flows.clientCredentials is not None + assert scheme.flows.clientCredentials.tokenUrl assert credential.auth_type == AuthCredentialTypes.SERVICE_ACCOUNT assert credential.service_account == config +@pytest.mark.asyncio +async def test_service_account_helper_scheme_allows_credential_manager_exchange(): + """SA helpers must yield a client-credentials scheme (#6656). + + With HTTPBearer, CredentialManager treated the SA as needing interactive + auth and returned None (adk_request_credential) instead of exchanging it. + """ + scheme, credential = service_account_scheme_credential( + ServiceAccount( + use_default_credential=True, + scopes=["https://www.googleapis.com/auth/cloud-platform"], + ) + ) + manager = CredentialManager( + AuthConfig(auth_scheme=scheme, raw_auth_credential=credential) + ) + assert manager._is_client_credentials_flow() # pylint: disable=protected-access + + exchanged = AuthCredential( + auth_type=AuthCredentialTypes.HTTP, + http=HttpAuth( + scheme="bearer", + credentials=HttpCredentials(token="sa-access-token"), + ), + ) + manager._load_existing_credential = AsyncMock(return_value=None) # pylint: disable=protected-access + manager._exchange_credential = AsyncMock(return_value=(exchanged, True)) # pylint: disable=protected-access + manager._refresh_credential = AsyncMock(return_value=(exchanged, False)) # pylint: disable=protected-access + manager._save_credential = AsyncMock() # pylint: disable=protected-access + + ctx = Mock() + ctx.get_auth_response = Mock(return_value=None) + result = await manager.get_auth_credential(ctx) + + assert result is not None + assert result.auth_type == AuthCredentialTypes.HTTP + assert result.http.credentials.token == "sa-access-token" + manager._exchange_credential.assert_awaited_once() # pylint: disable=protected-access + + def test_openid_dict_to_scheme_credential(): config_dict = { "authorization_endpoint": "auth_url",