mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
refactor: Refactor oauth2_credential_exchanger to exchanger and refresher separately
PiperOrigin-RevId: 772979993
This commit is contained in:
committed by
Copybara-Service
parent
a17ebe6ebd
commit
9a207cb832
@@ -22,7 +22,7 @@ from .auth_credential import AuthCredential
|
||||
from .auth_schemes import AuthSchemeType
|
||||
from .auth_schemes import OpenIdConnectWithConfig
|
||||
from .auth_tool import AuthConfig
|
||||
from .oauth2_credential_fetcher import OAuth2CredentialFetcher
|
||||
from .exchanger.oauth2_credential_exchanger import OAuth2CredentialExchanger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..sessions.state import State
|
||||
@@ -36,18 +36,23 @@ except ImportError:
|
||||
|
||||
|
||||
class AuthHandler:
|
||||
"""A handler that handles the auth flow in Agent Development Kit to help
|
||||
orchestrate the credential request and response flow (e.g. OAuth flow)
|
||||
This class should only be used by Agent Development Kit.
|
||||
"""
|
||||
|
||||
def __init__(self, auth_config: AuthConfig):
|
||||
self.auth_config = auth_config
|
||||
|
||||
def exchange_auth_token(
|
||||
async def exchange_auth_token(
|
||||
self,
|
||||
) -> AuthCredential:
|
||||
return OAuth2CredentialFetcher(
|
||||
self.auth_config.auth_scheme, self.auth_config.exchanged_auth_credential
|
||||
).exchange()
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
return await exchanger.exchange(
|
||||
self.auth_config.exchanged_auth_credential, self.auth_config.auth_scheme
|
||||
)
|
||||
|
||||
def parse_and_store_auth_response(self, state: State) -> None:
|
||||
async def parse_and_store_auth_response(self, state: State) -> None:
|
||||
|
||||
credential_key = "temp:" + self.auth_config.credential_key
|
||||
|
||||
@@ -60,7 +65,7 @@ class AuthHandler:
|
||||
):
|
||||
return
|
||||
|
||||
state[credential_key] = self.exchange_auth_token()
|
||||
state[credential_key] = await self.exchange_auth_token()
|
||||
|
||||
def _validate(self) -> None:
|
||||
if not self.auth_scheme:
|
||||
|
||||
@@ -67,9 +67,9 @@ class _AuthLlmRequestProcessor(BaseLlmRequestProcessor):
|
||||
# function call
|
||||
request_euc_function_call_ids.add(function_call_response.id)
|
||||
auth_config = AuthConfig.model_validate(function_call_response.response)
|
||||
AuthHandler(auth_config=auth_config).parse_and_store_auth_response(
|
||||
state=invocation_context.session.state
|
||||
)
|
||||
await AuthHandler(
|
||||
auth_config=auth_config
|
||||
).parse_and_store_auth_response(state=invocation_context.session.state)
|
||||
break
|
||||
|
||||
if not request_euc_function_call_ids:
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
# Copyright 2025 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.
|
||||
|
||||
"""OAuth2 credential exchanger implementation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from google.adk.auth.auth_credential import AuthCredential
|
||||
from google.adk.auth.auth_schemes import AuthScheme
|
||||
from google.adk.auth.auth_schemes import OAuthGrantType
|
||||
from google.adk.auth.oauth2_credential_util import create_oauth2_session
|
||||
from google.adk.auth.oauth2_credential_util import update_credential_with_tokens
|
||||
from google.adk.utils.feature_decorator import experimental
|
||||
from typing_extensions import override
|
||||
|
||||
from .base_credential_exchanger import BaseCredentialExchanger
|
||||
from .base_credential_exchanger import CredentialExchangError
|
||||
|
||||
try:
|
||||
from authlib.integrations.requests_client import OAuth2Session
|
||||
|
||||
AUTHLIB_AVIALABLE = True
|
||||
except ImportError:
|
||||
AUTHLIB_AVIALABLE = False
|
||||
|
||||
logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
|
||||
@experimental
|
||||
class OAuth2CredentialExchanger(BaseCredentialExchanger):
|
||||
"""Exchanges OAuth2 credentials from authorization responses."""
|
||||
|
||||
@override
|
||||
async def exchange(
|
||||
self,
|
||||
auth_credential: AuthCredential,
|
||||
auth_scheme: Optional[AuthScheme] = None,
|
||||
) -> AuthCredential:
|
||||
"""Exchange OAuth2 credential from authorization response.
|
||||
if credential exchange failed, the original credential will be returned.
|
||||
|
||||
Args:
|
||||
auth_credential: The OAuth2 credential to exchange.
|
||||
auth_scheme: The OAuth2 authentication scheme.
|
||||
|
||||
Returns:
|
||||
The exchanged credential with access token.
|
||||
|
||||
Raises:
|
||||
CredentialExchangError: If auth_scheme is missing.
|
||||
"""
|
||||
if not auth_scheme:
|
||||
raise CredentialExchangError(
|
||||
"auth_scheme is required for OAuth2 credential exchange"
|
||||
)
|
||||
|
||||
if not AUTHLIB_AVIALABLE:
|
||||
# If authlib is not available, we cannot exchange the credential.
|
||||
# We return the original credential without exchange.
|
||||
# The client using this tool can decide to exchange the credential
|
||||
# themselves using other lib.
|
||||
logger.warning(
|
||||
"authlib is not available, skipping OAuth2 credential exchange."
|
||||
)
|
||||
return auth_credential
|
||||
|
||||
if auth_credential.oauth2 and auth_credential.oauth2.access_token:
|
||||
return auth_credential
|
||||
|
||||
client, token_endpoint = create_oauth2_session(auth_scheme, auth_credential)
|
||||
if not client:
|
||||
logger.warning("Could not create OAuth2 session for token exchange")
|
||||
return auth_credential
|
||||
|
||||
try:
|
||||
tokens = client.fetch_token(
|
||||
token_endpoint,
|
||||
authorization_response=auth_credential.oauth2.auth_response_uri,
|
||||
code=auth_credential.oauth2.auth_code,
|
||||
grant_type=OAuthGrantType.AUTHORIZATION_CODE,
|
||||
)
|
||||
update_credential_with_tokens(auth_credential, tokens)
|
||||
logger.debug("Successfully exchanged OAuth2 tokens")
|
||||
except Exception as e:
|
||||
# TODO reconsider whether we should raise errors in this case
|
||||
logger.error("Failed to exchange OAuth2 tokens: %s", e)
|
||||
# Return original credential on failure
|
||||
return auth_credential
|
||||
|
||||
return auth_credential
|
||||
@@ -1,132 +0,0 @@
|
||||
# Copyright 2025 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 logging
|
||||
|
||||
from ..utils.feature_decorator import experimental
|
||||
from .auth_credential import AuthCredential
|
||||
from .auth_schemes import AuthScheme
|
||||
from .auth_schemes import OAuthGrantType
|
||||
from .oauth2_credential_util import create_oauth2_session
|
||||
from .oauth2_credential_util import update_credential_with_tokens
|
||||
|
||||
try:
|
||||
from authlib.oauth2.rfc6749 import OAuth2Token
|
||||
|
||||
AUTHLIB_AVIALABLE = True
|
||||
except ImportError:
|
||||
AUTHLIB_AVIALABLE = False
|
||||
|
||||
|
||||
logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
|
||||
@experimental
|
||||
class OAuth2CredentialFetcher:
|
||||
"""Exchanges and refreshes an OAuth2 access token. (Experimental)"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
auth_scheme: AuthScheme,
|
||||
auth_credential: AuthCredential,
|
||||
):
|
||||
self._auth_scheme = auth_scheme
|
||||
self._auth_credential = auth_credential
|
||||
|
||||
def _update_credential(self, tokens: OAuth2Token) -> None:
|
||||
self._auth_credential.oauth2.access_token = tokens.get("access_token")
|
||||
self._auth_credential.oauth2.refresh_token = tokens.get("refresh_token")
|
||||
self._auth_credential.oauth2.expires_at = (
|
||||
int(tokens.get("expires_at")) if tokens.get("expires_at") else None
|
||||
)
|
||||
self._auth_credential.oauth2.expires_in = (
|
||||
int(tokens.get("expires_in")) if tokens.get("expires_in") else None
|
||||
)
|
||||
|
||||
def exchange(self) -> AuthCredential:
|
||||
"""Exchange an oauth token from the authorization response.
|
||||
|
||||
Returns:
|
||||
An AuthCredential object containing the access token.
|
||||
"""
|
||||
if not AUTHLIB_AVIALABLE:
|
||||
return self._auth_credential
|
||||
|
||||
if (
|
||||
self._auth_credential.oauth2
|
||||
and self._auth_credential.oauth2.access_token
|
||||
):
|
||||
return self._auth_credential
|
||||
|
||||
client, token_endpoint = create_oauth2_session(
|
||||
self._auth_scheme, self._auth_credential
|
||||
)
|
||||
if not client:
|
||||
logger.warning("Could not create OAuth2 session for token exchange")
|
||||
return self._auth_credential
|
||||
|
||||
try:
|
||||
tokens = client.fetch_token(
|
||||
token_endpoint,
|
||||
authorization_response=self._auth_credential.oauth2.auth_response_uri,
|
||||
code=self._auth_credential.oauth2.auth_code,
|
||||
grant_type=OAuthGrantType.AUTHORIZATION_CODE,
|
||||
)
|
||||
update_credential_with_tokens(self._auth_credential, tokens)
|
||||
logger.info("Successfully exchanged OAuth2 tokens")
|
||||
except Exception as e:
|
||||
logger.error("Failed to exchange OAuth2 tokens: %s", e)
|
||||
# Return original credential on failure
|
||||
return self._auth_credential
|
||||
|
||||
return self._auth_credential
|
||||
|
||||
def refresh(self) -> AuthCredential:
|
||||
"""Refresh an oauth token.
|
||||
|
||||
Returns:
|
||||
An AuthCredential object containing the refreshed access token.
|
||||
"""
|
||||
if not AUTHLIB_AVIALABLE:
|
||||
return self._auth_credential
|
||||
credential = self._auth_credential
|
||||
if not credential.oauth2:
|
||||
return credential
|
||||
|
||||
if OAuth2Token({
|
||||
"expires_at": credential.oauth2.expires_at,
|
||||
"expires_in": credential.oauth2.expires_in,
|
||||
}).is_expired():
|
||||
client, token_endpoint = create_oauth2_session(
|
||||
self._auth_scheme, self._auth_credential
|
||||
)
|
||||
if not client:
|
||||
logger.warning("Could not create OAuth2 session for token refresh")
|
||||
return credential
|
||||
|
||||
try:
|
||||
tokens = client.refresh_token(
|
||||
url=token_endpoint,
|
||||
refresh_token=credential.oauth2.refresh_token,
|
||||
)
|
||||
update_credential_with_tokens(self._auth_credential, tokens)
|
||||
logger.info("Successfully refreshed OAuth2 tokens")
|
||||
except Exception as e:
|
||||
logger.error("Failed to refresh OAuth2 tokens: %s", e)
|
||||
# Return original credential on failure
|
||||
return credential
|
||||
|
||||
return self._auth_credential
|
||||
@@ -0,0 +1,154 @@
|
||||
# Copyright 2025 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.
|
||||
|
||||
"""OAuth2 credential refresher implementation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from google.adk.auth.auth_credential import AuthCredential
|
||||
from google.adk.auth.auth_schemes import AuthScheme
|
||||
from google.adk.auth.oauth2_credential_util import create_oauth2_session
|
||||
from google.adk.auth.oauth2_credential_util import update_credential_with_tokens
|
||||
from google.adk.utils.feature_decorator import experimental
|
||||
from google.auth.transport.requests import Request
|
||||
from google.oauth2.credentials import Credentials
|
||||
from typing_extensions import override
|
||||
|
||||
from .base_credential_refresher import BaseCredentialRefresher
|
||||
|
||||
try:
|
||||
from authlib.oauth2.rfc6749 import OAuth2Token
|
||||
|
||||
AUTHLIB_AVIALABLE = True
|
||||
except ImportError:
|
||||
AUTHLIB_AVIALABLE = False
|
||||
|
||||
logger = logging.getLogger("google_adk." + __name__)
|
||||
|
||||
|
||||
@experimental
|
||||
class OAuth2CredentialRefresher(BaseCredentialRefresher):
|
||||
"""Refreshes OAuth2 credentials including Google OAuth2 JSON credentials."""
|
||||
|
||||
@override
|
||||
async def is_refresh_needed(
|
||||
self,
|
||||
auth_credential: AuthCredential,
|
||||
auth_scheme: Optional[AuthScheme] = None,
|
||||
) -> bool:
|
||||
"""Check if the OAuth2 credential needs to be refreshed.
|
||||
|
||||
Args:
|
||||
auth_credential: The OAuth2 credential to check.
|
||||
auth_scheme: The OAuth2 authentication scheme (optional for Google OAuth2 JSON).
|
||||
|
||||
Returns:
|
||||
True if the credential needs to be refreshed, False otherwise.
|
||||
"""
|
||||
# Handle Google OAuth2 credentials (from service account exchange)
|
||||
if auth_credential.google_oauth2_json:
|
||||
try:
|
||||
google_credential = Credentials.from_authorized_user_info(
|
||||
json.loads(auth_credential.google_oauth2_json)
|
||||
)
|
||||
return google_credential.expired and bool(
|
||||
google_credential.refresh_token
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to parse Google OAuth2 JSON credential: %s", e)
|
||||
return False
|
||||
|
||||
# Handle regular OAuth2 credentials
|
||||
elif auth_credential.oauth2 and auth_scheme:
|
||||
if not AUTHLIB_AVIALABLE:
|
||||
return False
|
||||
|
||||
if not auth_credential.oauth2:
|
||||
return False
|
||||
|
||||
return OAuth2Token({
|
||||
"expires_at": auth_credential.oauth2.expires_at,
|
||||
"expires_in": auth_credential.oauth2.expires_in,
|
||||
}).is_expired()
|
||||
|
||||
return False
|
||||
|
||||
@override
|
||||
async def refresh(
|
||||
self,
|
||||
auth_credential: AuthCredential,
|
||||
auth_scheme: Optional[AuthScheme] = None,
|
||||
) -> AuthCredential:
|
||||
"""Refresh the OAuth2 credential.
|
||||
If refresh failed, return the original credential.
|
||||
|
||||
Args:
|
||||
auth_credential: The OAuth2 credential to refresh.
|
||||
auth_scheme: The OAuth2 authentication scheme (optional for Google OAuth2 JSON).
|
||||
|
||||
Returns:
|
||||
The refreshed credential.
|
||||
|
||||
"""
|
||||
# Handle Google OAuth2 credentials (from service account exchange)
|
||||
if auth_credential.google_oauth2_json:
|
||||
try:
|
||||
google_credential = Credentials.from_authorized_user_info(
|
||||
json.loads(auth_credential.google_oauth2_json)
|
||||
)
|
||||
if google_credential.expired and google_credential.refresh_token:
|
||||
google_credential.refresh(Request())
|
||||
auth_credential.google_oauth2_json = google_credential.to_json()
|
||||
logger.info("Successfully refreshed Google OAuth2 JSON credential")
|
||||
except Exception as e:
|
||||
# TODO reconsider whether we should raise error when refresh failed.
|
||||
logger.error("Failed to refresh Google OAuth2 JSON credential: %s", e)
|
||||
|
||||
# Handle regular OAuth2 credentials
|
||||
elif auth_credential.oauth2 and auth_scheme:
|
||||
if not AUTHLIB_AVIALABLE:
|
||||
return auth_credential
|
||||
|
||||
if not auth_credential.oauth2:
|
||||
return auth_credential
|
||||
|
||||
if OAuth2Token({
|
||||
"expires_at": auth_credential.oauth2.expires_at,
|
||||
"expires_in": auth_credential.oauth2.expires_in,
|
||||
}).is_expired():
|
||||
client, token_endpoint = create_oauth2_session(
|
||||
auth_scheme, auth_credential
|
||||
)
|
||||
if not client:
|
||||
logger.warning("Could not create OAuth2 session for token refresh")
|
||||
return auth_credential
|
||||
|
||||
try:
|
||||
tokens = client.refresh_token(
|
||||
url=token_endpoint,
|
||||
refresh_token=auth_credential.oauth2.refresh_token,
|
||||
)
|
||||
update_credential_with_tokens(auth_credential, tokens)
|
||||
logger.debug("Successfully refreshed OAuth2 tokens")
|
||||
except Exception as e:
|
||||
# TODO reconsider whether we should raise error when refresh failed.
|
||||
logger.error("Failed to refresh OAuth2 tokens: %s", e)
|
||||
# Return original credential on failure
|
||||
return auth_credential
|
||||
|
||||
return auth_credential
|
||||
@@ -150,7 +150,7 @@ class IntegrationConnectorTool(BaseTool):
|
||||
tool_auth_handler = ToolAuthHandler.from_tool_context(
|
||||
tool_context, self._auth_scheme, self._auth_credential
|
||||
)
|
||||
auth_result = tool_auth_handler.prepare_auth_credentials()
|
||||
auth_result = await tool_auth_handler.prepare_auth_credentials()
|
||||
|
||||
if auth_result.state == 'pending':
|
||||
return {
|
||||
@@ -178,7 +178,7 @@ class IntegrationConnectorTool(BaseTool):
|
||||
args['operation'] = self._operation
|
||||
args['action'] = self._action
|
||||
logger.info('Running tool: %s with args: %s', self.name, args)
|
||||
return self._rest_api_tool.call(args=args, tool_context=tool_context)
|
||||
return await self._rest_api_tool.call(args=args, tool_context=tool_context)
|
||||
|
||||
def __str__(self):
|
||||
return (
|
||||
|
||||
@@ -345,9 +345,9 @@ class RestApiTool(BaseTool):
|
||||
async def run_async(
|
||||
self, *, args: dict[str, Any], tool_context: Optional[ToolContext]
|
||||
) -> Dict[str, Any]:
|
||||
return self.call(args=args, tool_context=tool_context)
|
||||
return await self.call(args=args, tool_context=tool_context)
|
||||
|
||||
def call(
|
||||
async def call(
|
||||
self, *, args: dict[str, Any], tool_context: Optional[ToolContext]
|
||||
) -> Dict[str, Any]:
|
||||
"""Executes the REST API call.
|
||||
@@ -364,7 +364,7 @@ class RestApiTool(BaseTool):
|
||||
tool_auth_handler = ToolAuthHandler.from_tool_context(
|
||||
tool_context, self.auth_scheme, self.auth_credential
|
||||
)
|
||||
auth_result = tool_auth_handler.prepare_auth_credentials()
|
||||
auth_result = await tool_auth_handler.prepare_auth_credentials()
|
||||
auth_state, auth_scheme, auth_credential = (
|
||||
auth_result.state,
|
||||
auth_result.auth_scheme,
|
||||
|
||||
@@ -25,7 +25,7 @@ from ....auth.auth_credential import AuthCredentialTypes
|
||||
from ....auth.auth_schemes import AuthScheme
|
||||
from ....auth.auth_schemes import AuthSchemeType
|
||||
from ....auth.auth_tool import AuthConfig
|
||||
from ....auth.oauth2_credential_fetcher import OAuth2CredentialFetcher
|
||||
from ....auth.refresher.oauth2_credential_refresher import OAuth2CredentialRefresher
|
||||
from ...tool_context import ToolContext
|
||||
from ..auth.credential_exchangers.auto_auth_credential_exchanger import AutoAuthCredentialExchanger
|
||||
from ..auth.credential_exchangers.base_credential_exchanger import AuthCredentialMissingError
|
||||
@@ -146,7 +146,7 @@ class ToolAuthHandler:
|
||||
credential_store,
|
||||
)
|
||||
|
||||
def _get_existing_credential(
|
||||
async def _get_existing_credential(
|
||||
self,
|
||||
) -> Optional[AuthCredential]:
|
||||
"""Checks for and returns an existing, exchanged credential."""
|
||||
@@ -156,9 +156,11 @@ class ToolAuthHandler:
|
||||
)
|
||||
if existing_credential:
|
||||
if existing_credential.oauth2:
|
||||
existing_credential = OAuth2CredentialFetcher(
|
||||
self.auth_scheme, existing_credential
|
||||
).refresh()
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
if await refresher.is_refresh_needed(existing_credential):
|
||||
existing_credential = await refresher.refresh(
|
||||
existing_credential, self.auth_scheme
|
||||
)
|
||||
return existing_credential
|
||||
return None
|
||||
|
||||
@@ -234,7 +236,7 @@ class ToolAuthHandler:
|
||||
and not credential.google_oauth2_json
|
||||
)
|
||||
|
||||
def prepare_auth_credentials(
|
||||
async def prepare_auth_credentials(
|
||||
self,
|
||||
) -> AuthPreparationResult:
|
||||
"""Prepares authentication credentials, handling exchange and user interaction."""
|
||||
@@ -244,7 +246,7 @@ class ToolAuthHandler:
|
||||
return AuthPreparationResult(state="done")
|
||||
|
||||
# Check for existing credential.
|
||||
existing_credential = self._get_existing_credential()
|
||||
existing_credential = await self._get_existing_credential()
|
||||
|
||||
credential = existing_credential or self.auth_credential
|
||||
# fetch credential from adk framework
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
# Copyright 2025 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.
|
||||
|
||||
import time
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import patch
|
||||
|
||||
from authlib.oauth2.rfc6749 import OAuth2Token
|
||||
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_schemes import OpenIdConnectWithConfig
|
||||
from google.adk.auth.exchanger.base_credential_exchanger import CredentialExchangError
|
||||
from google.adk.auth.exchanger.oauth2_credential_exchanger import OAuth2CredentialExchanger
|
||||
import pytest
|
||||
|
||||
|
||||
class TestOAuth2CredentialExchanger:
|
||||
"""Test suite for OAuth2CredentialExchanger."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_with_existing_token(self):
|
||||
"""Test exchange method when access token already exists."""
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
access_token="existing_token",
|
||||
),
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Should return the same credential since access token already exists
|
||||
assert result == credential
|
||||
assert result.oauth2.access_token == "existing_token"
|
||||
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Session")
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_success(self, mock_oauth2_session):
|
||||
"""Test successful token exchange."""
|
||||
# Setup mock
|
||||
mock_client = Mock()
|
||||
mock_oauth2_session.return_value = mock_client
|
||||
mock_tokens = OAuth2Token({
|
||||
"access_token": "new_access_token",
|
||||
"refresh_token": "new_refresh_token",
|
||||
"expires_at": int(time.time()) + 3600,
|
||||
"expires_in": 3600,
|
||||
})
|
||||
mock_client.fetch_token.return_value = mock_tokens
|
||||
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
auth_response_uri="https://example.com/callback?code=auth_code",
|
||||
auth_code="auth_code",
|
||||
),
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Verify token exchange was successful
|
||||
assert result.oauth2.access_token == "new_access_token"
|
||||
assert result.oauth2.refresh_token == "new_refresh_token"
|
||||
mock_client.fetch_token.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_missing_auth_scheme(self):
|
||||
"""Test exchange with missing auth_scheme raises ValueError."""
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
),
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
try:
|
||||
await exchanger.exchange(credential, None)
|
||||
assert False, "Should have raised ValueError"
|
||||
except CredentialExchangError as e:
|
||||
assert "auth_scheme is required" in str(e)
|
||||
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Session")
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_no_session(self, mock_oauth2_session):
|
||||
"""Test exchange when OAuth2Session cannot be created."""
|
||||
# Mock to return None for create_oauth2_session
|
||||
mock_oauth2_session.return_value = None
|
||||
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
# Missing client_secret to trigger session creation failure
|
||||
),
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Should return original credential when session creation fails
|
||||
assert result == credential
|
||||
assert result.oauth2.access_token is None
|
||||
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Session")
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_fetch_token_failure(self, mock_oauth2_session):
|
||||
"""Test exchange when fetch_token fails."""
|
||||
# Setup mock to raise exception during fetch_token
|
||||
mock_client = Mock()
|
||||
mock_oauth2_session.return_value = mock_client
|
||||
mock_client.fetch_token.side_effect = Exception("Token fetch failed")
|
||||
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
auth_response_uri="https://example.com/callback?code=auth_code",
|
||||
auth_code="auth_code",
|
||||
),
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Should return original credential when fetch_token fails
|
||||
assert result == credential
|
||||
assert result.oauth2.access_token is None
|
||||
mock_client.fetch_token.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_authlib_not_available(self):
|
||||
"""Test exchange when authlib is not available."""
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
auth_response_uri="https://example.com/callback?code=auth_code",
|
||||
auth_code="auth_code",
|
||||
),
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
|
||||
# Mock AUTHLIB_AVIALABLE to False
|
||||
with patch(
|
||||
"google.adk.auth.exchanger.oauth2_credential_exchanger.AUTHLIB_AVIALABLE",
|
||||
False,
|
||||
):
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Should return original credential when authlib is not available
|
||||
assert result == credential
|
||||
assert result.oauth2.access_token is None
|
||||
@@ -0,0 +1,13 @@
|
||||
# Copyright 2025 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.
|
||||
@@ -0,0 +1,297 @@
|
||||
# Copyright 2025 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.
|
||||
|
||||
import time
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import patch
|
||||
|
||||
from authlib.oauth2.rfc6749 import OAuth2Token
|
||||
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_schemes import OpenIdConnectWithConfig
|
||||
from google.adk.auth.refresher.oauth2_credential_refresher import OAuth2CredentialRefresher
|
||||
import pytest
|
||||
|
||||
|
||||
class TestOAuth2CredentialRefresher:
|
||||
"""Test suite for OAuth2CredentialRefresher."""
|
||||
|
||||
@patch("google.adk.auth.refresher.oauth2_credential_refresher.OAuth2Token")
|
||||
@pytest.mark.asyncio
|
||||
async def test_needs_refresh_token_not_expired(self, mock_oauth2_token):
|
||||
"""Test needs_refresh when token is not expired."""
|
||||
mock_token_instance = Mock()
|
||||
mock_token_instance.is_expired.return_value = False
|
||||
mock_oauth2_token.return_value = mock_token_instance
|
||||
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
access_token="existing_token",
|
||||
expires_at=int(time.time()) + 3600,
|
||||
),
|
||||
)
|
||||
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
needs_refresh = await refresher.is_refresh_needed(credential, scheme)
|
||||
|
||||
assert not needs_refresh
|
||||
|
||||
@patch("google.adk.auth.refresher.oauth2_credential_refresher.OAuth2Token")
|
||||
@pytest.mark.asyncio
|
||||
async def test_needs_refresh_token_expired(self, mock_oauth2_token):
|
||||
"""Test needs_refresh when token is expired."""
|
||||
mock_token_instance = Mock()
|
||||
mock_token_instance.is_expired.return_value = True
|
||||
mock_oauth2_token.return_value = mock_token_instance
|
||||
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
access_token="existing_token",
|
||||
expires_at=int(time.time()) - 3600, # Expired
|
||||
),
|
||||
)
|
||||
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
needs_refresh = await refresher.is_refresh_needed(credential, scheme)
|
||||
|
||||
assert needs_refresh
|
||||
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Session")
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Token")
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_token_expired_success(
|
||||
self, mock_oauth2_token, mock_oauth2_session
|
||||
):
|
||||
"""Test successful token refresh when token is expired."""
|
||||
# Setup mock token
|
||||
mock_token_instance = Mock()
|
||||
mock_token_instance.is_expired.return_value = True
|
||||
mock_oauth2_token.return_value = mock_token_instance
|
||||
|
||||
# Setup mock session
|
||||
mock_client = Mock()
|
||||
mock_oauth2_session.return_value = mock_client
|
||||
mock_tokens = OAuth2Token({
|
||||
"access_token": "refreshed_access_token",
|
||||
"refresh_token": "refreshed_refresh_token",
|
||||
"expires_at": int(time.time()) + 3600,
|
||||
"expires_in": 3600,
|
||||
})
|
||||
mock_client.refresh_token.return_value = mock_tokens
|
||||
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
access_token="old_token",
|
||||
refresh_token="old_refresh_token",
|
||||
expires_at=int(time.time()) - 3600, # Expired
|
||||
),
|
||||
)
|
||||
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
result = await refresher.refresh(credential, scheme)
|
||||
|
||||
# Verify token refresh was successful
|
||||
assert result.oauth2.access_token == "refreshed_access_token"
|
||||
assert result.oauth2.refresh_token == "refreshed_refresh_token"
|
||||
mock_client.refresh_token.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_no_oauth2_credential(self):
|
||||
"""Test refresh with no OAuth2 credential returns original."""
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
# No oauth2 field
|
||||
)
|
||||
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
result = await refresher.refresh(credential, scheme)
|
||||
|
||||
assert result == credential
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_needs_refresh_google_oauth2_json_expired(self):
|
||||
"""Test needs_refresh with Google OAuth2 JSON credential that is expired."""
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
# Mock Google OAuth2 JSON credential data
|
||||
google_oauth2_json = json.dumps({
|
||||
"client_id": "test_client_id",
|
||||
"client_secret": "test_client_secret",
|
||||
"refresh_token": "test_refresh_token",
|
||||
"type": "authorized_user",
|
||||
})
|
||||
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
google_oauth2_json=google_oauth2_json,
|
||||
)
|
||||
|
||||
# Mock the Google Credentials class
|
||||
with patch(
|
||||
"google.adk.auth.refresher.oauth2_credential_refresher.Credentials"
|
||||
) as mock_credentials:
|
||||
mock_google_credential = Mock()
|
||||
mock_google_credential.expired = True
|
||||
mock_google_credential.refresh_token = "test_refresh_token"
|
||||
mock_credentials.from_authorized_user_info.return_value = (
|
||||
mock_google_credential
|
||||
)
|
||||
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
needs_refresh = await refresher.is_refresh_needed(credential, None)
|
||||
|
||||
assert needs_refresh
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_needs_refresh_google_oauth2_json_not_expired(self):
|
||||
"""Test needs_refresh with Google OAuth2 JSON credential that is not expired."""
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
# Mock Google OAuth2 JSON credential data
|
||||
google_oauth2_json = json.dumps({
|
||||
"client_id": "test_client_id",
|
||||
"client_secret": "test_client_secret",
|
||||
"refresh_token": "test_refresh_token",
|
||||
"type": "authorized_user",
|
||||
})
|
||||
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
google_oauth2_json=google_oauth2_json,
|
||||
)
|
||||
|
||||
# Mock the Google Credentials class
|
||||
with patch(
|
||||
"google.adk.auth.refresher.oauth2_credential_refresher.Credentials"
|
||||
) as mock_credentials:
|
||||
mock_google_credential = Mock()
|
||||
mock_google_credential.expired = False
|
||||
mock_google_credential.refresh_token = "test_refresh_token"
|
||||
mock_credentials.from_authorized_user_info.return_value = (
|
||||
mock_google_credential
|
||||
)
|
||||
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
needs_refresh = await refresher.is_refresh_needed(credential, None)
|
||||
|
||||
assert not needs_refresh
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_google_oauth2_json_success(self):
|
||||
"""Test successful refresh of Google OAuth2 JSON credential."""
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
# Mock Google OAuth2 JSON credential data
|
||||
google_oauth2_json = json.dumps({
|
||||
"client_id": "test_client_id",
|
||||
"client_secret": "test_client_secret",
|
||||
"refresh_token": "test_refresh_token",
|
||||
"type": "authorized_user",
|
||||
})
|
||||
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
google_oauth2_json=google_oauth2_json,
|
||||
)
|
||||
|
||||
# Mock the Google Credentials and Request classes
|
||||
with patch(
|
||||
"google.adk.auth.refresher.oauth2_credential_refresher.Credentials"
|
||||
) as mock_credentials:
|
||||
with patch(
|
||||
"google.adk.auth.refresher.oauth2_credential_refresher.Request"
|
||||
) as mock_request:
|
||||
mock_google_credential = Mock()
|
||||
mock_google_credential.expired = True
|
||||
mock_google_credential.refresh_token = "test_refresh_token"
|
||||
mock_google_credential.to_json.return_value = json.dumps({
|
||||
"client_id": "test_client_id",
|
||||
"client_secret": "test_client_secret",
|
||||
"refresh_token": "new_refresh_token",
|
||||
"access_token": "new_access_token",
|
||||
"type": "authorized_user",
|
||||
})
|
||||
mock_credentials.from_authorized_user_info.return_value = (
|
||||
mock_google_credential
|
||||
)
|
||||
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
result = await refresher.refresh(credential, None)
|
||||
|
||||
mock_google_credential.refresh.assert_called_once()
|
||||
assert (
|
||||
result.google_oauth2_json != google_oauth2_json
|
||||
) # Should be updated
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_needs_refresh_no_oauth2_credential(self):
|
||||
"""Test needs_refresh with no OAuth2 credential returns False."""
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.HTTP,
|
||||
# No oauth2 field
|
||||
)
|
||||
|
||||
refresher = OAuth2CredentialRefresher()
|
||||
needs_refresh = await refresher.is_refresh_needed(credential, None)
|
||||
|
||||
assert not needs_refresh
|
||||
@@ -13,8 +13,11 @@
|
||||
# limitations under the License.
|
||||
|
||||
import copy
|
||||
import time
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import patch
|
||||
|
||||
from authlib.oauth2.rfc6749 import OAuth2Token
|
||||
from fastapi.openapi.models import APIKey
|
||||
from fastapi.openapi.models import APIKeyIn
|
||||
from fastapi.openapi.models import OAuth2
|
||||
@@ -405,7 +408,8 @@ class TestGetAuthResponse:
|
||||
class TestParseAndStoreAuthResponse:
|
||||
"""Tests for the parse_and_store_auth_response method."""
|
||||
|
||||
def test_non_oauth_scheme(self, auth_config_with_exchanged):
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_oauth_scheme(self, auth_config_with_exchanged):
|
||||
"""Test with a non-OAuth auth scheme."""
|
||||
# Modify the auth scheme type to be non-OAuth
|
||||
auth_config = copy.deepcopy(auth_config_with_exchanged)
|
||||
@@ -416,7 +420,7 @@ class TestParseAndStoreAuthResponse:
|
||||
handler = AuthHandler(auth_config)
|
||||
state = MockState()
|
||||
|
||||
handler.parse_and_store_auth_response(state)
|
||||
await handler.parse_and_store_auth_response(state)
|
||||
|
||||
credential_key = auth_config.credential_key
|
||||
assert (
|
||||
@@ -424,7 +428,10 @@ class TestParseAndStoreAuthResponse:
|
||||
)
|
||||
|
||||
@patch("google.adk.auth.auth_handler.AuthHandler.exchange_auth_token")
|
||||
def test_oauth_scheme(self, mock_exchange_token, auth_config_with_exchanged):
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_scheme(
|
||||
self, mock_exchange_token, auth_config_with_exchanged
|
||||
):
|
||||
"""Test with an OAuth auth scheme."""
|
||||
mock_exchange_token.return_value = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
@@ -434,7 +441,7 @@ class TestParseAndStoreAuthResponse:
|
||||
handler = AuthHandler(auth_config_with_exchanged)
|
||||
state = MockState()
|
||||
|
||||
handler.parse_and_store_auth_response(state)
|
||||
await handler.parse_and_store_auth_response(state)
|
||||
|
||||
credential_key = auth_config_with_exchanged.credential_key
|
||||
assert state["temp:" + credential_key] == mock_exchange_token.return_value
|
||||
@@ -444,20 +451,20 @@ class TestParseAndStoreAuthResponse:
|
||||
class TestExchangeAuthToken:
|
||||
"""Tests for the exchange_auth_token method."""
|
||||
|
||||
def test_token_exchange_not_supported(
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_not_supported(
|
||||
self, auth_config_with_auth_code, monkeypatch
|
||||
):
|
||||
"""Test when token exchange is not supported."""
|
||||
monkeypatch.setattr(
|
||||
"google.adk.auth.oauth2_credential_fetcher.AUTHLIB_AVIALABLE", False
|
||||
)
|
||||
monkeypatch.setattr("google.adk.auth.auth_handler.AUTHLIB_AVIALABLE", False)
|
||||
|
||||
handler = AuthHandler(auth_config_with_auth_code)
|
||||
result = handler.exchange_auth_token()
|
||||
result = await handler.exchange_auth_token()
|
||||
|
||||
assert result == auth_config_with_auth_code.exchanged_auth_credential
|
||||
|
||||
def test_openid_missing_token_endpoint(
|
||||
@pytest.mark.asyncio
|
||||
async def test_openid_missing_token_endpoint(
|
||||
self, openid_auth_scheme, oauth2_credentials_with_auth_code
|
||||
):
|
||||
"""Test OpenID Connect without a token endpoint."""
|
||||
@@ -472,11 +479,12 @@ class TestExchangeAuthToken:
|
||||
)
|
||||
|
||||
handler = AuthHandler(config)
|
||||
result = handler.exchange_auth_token()
|
||||
result = await handler.exchange_auth_token()
|
||||
|
||||
assert result == oauth2_credentials_with_auth_code
|
||||
|
||||
def test_oauth2_missing_token_url(
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth2_missing_token_url(
|
||||
self, oauth2_auth_scheme, oauth2_credentials_with_auth_code
|
||||
):
|
||||
"""Test OAuth2 without a token URL."""
|
||||
@@ -491,11 +499,12 @@ class TestExchangeAuthToken:
|
||||
)
|
||||
|
||||
handler = AuthHandler(config)
|
||||
result = handler.exchange_auth_token()
|
||||
result = await handler.exchange_auth_token()
|
||||
|
||||
assert result == oauth2_credentials_with_auth_code
|
||||
|
||||
def test_non_oauth_scheme(self, auth_config_with_auth_code):
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_oauth_scheme(self, auth_config_with_auth_code):
|
||||
"""Test with a non-OAuth auth scheme."""
|
||||
# Modify the auth scheme type to be non-OAuth
|
||||
auth_config = copy.deepcopy(auth_config_with_auth_code)
|
||||
@@ -504,11 +513,12 @@ class TestExchangeAuthToken:
|
||||
)
|
||||
|
||||
handler = AuthHandler(auth_config)
|
||||
result = handler.exchange_auth_token()
|
||||
result = await handler.exchange_auth_token()
|
||||
|
||||
assert result == auth_config.exchanged_auth_credential
|
||||
|
||||
def test_missing_credentials(self, oauth2_auth_scheme):
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_credentials(self, oauth2_auth_scheme):
|
||||
"""Test with missing credentials."""
|
||||
empty_credential = AuthCredential(auth_type=AuthCredentialTypes.OAUTH2)
|
||||
|
||||
@@ -518,11 +528,12 @@ class TestExchangeAuthToken:
|
||||
)
|
||||
|
||||
handler = AuthHandler(config)
|
||||
result = handler.exchange_auth_token()
|
||||
result = await handler.exchange_auth_token()
|
||||
|
||||
assert result == empty_credential
|
||||
|
||||
def test_credentials_with_token(
|
||||
@pytest.mark.asyncio
|
||||
async def test_credentials_with_token(
|
||||
self, auth_config, oauth2_credentials_with_token
|
||||
):
|
||||
"""Test when credentials already have a token."""
|
||||
@@ -533,18 +544,29 @@ class TestExchangeAuthToken:
|
||||
)
|
||||
|
||||
handler = AuthHandler(config)
|
||||
result = handler.exchange_auth_token()
|
||||
result = await handler.exchange_auth_token()
|
||||
|
||||
assert result == oauth2_credentials_with_token
|
||||
|
||||
@patch(
|
||||
"google.adk.auth.oauth2_credential_util.OAuth2Session",
|
||||
MockOAuth2Session,
|
||||
)
|
||||
def test_successful_token_exchange(self, auth_config_with_auth_code):
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Session")
|
||||
@pytest.mark.asyncio
|
||||
async def test_successful_token_exchange(
|
||||
self, mock_oauth2_session, auth_config_with_auth_code
|
||||
):
|
||||
"""Test a successful token exchange."""
|
||||
# Setup mock OAuth2Session
|
||||
mock_client = Mock()
|
||||
mock_oauth2_session.return_value = mock_client
|
||||
mock_tokens = OAuth2Token({
|
||||
"access_token": "mock_access_token",
|
||||
"refresh_token": "mock_refresh_token",
|
||||
"expires_at": int(time.time()) + 3600,
|
||||
"expires_in": 3600,
|
||||
})
|
||||
mock_client.fetch_token.return_value = mock_tokens
|
||||
|
||||
handler = AuthHandler(auth_config_with_auth_code)
|
||||
result = handler.exchange_auth_token()
|
||||
result = await handler.exchange_auth_token()
|
||||
|
||||
assert result.oauth2.access_token == "mock_access_token"
|
||||
assert result.oauth2.refresh_token == "mock_refresh_token"
|
||||
|
||||
@@ -1,147 +0,0 @@
|
||||
# Copyright 2025 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.
|
||||
|
||||
import time
|
||||
from unittest.mock import Mock
|
||||
|
||||
from authlib.oauth2.rfc6749 import OAuth2Token
|
||||
from fastapi.openapi.models import OAuth2
|
||||
from fastapi.openapi.models import OAuthFlowAuthorizationCode
|
||||
from fastapi.openapi.models import OAuthFlows
|
||||
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_schemes import OpenIdConnectWithConfig
|
||||
from google.adk.auth.oauth2_credential_util import create_oauth2_session
|
||||
from google.adk.auth.oauth2_credential_util import update_credential_with_tokens
|
||||
|
||||
|
||||
class TestOAuth2CredentialUtil:
|
||||
"""Test suite for OAuth2 credential utility functions."""
|
||||
|
||||
def test_create_oauth2_session_openid_connect(self):
|
||||
"""Test create_oauth2_session with OpenID Connect scheme."""
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid", "profile"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
redirect_uri="https://example.com/callback",
|
||||
state="test_state",
|
||||
),
|
||||
)
|
||||
|
||||
client, token_endpoint = create_oauth2_session(scheme, credential)
|
||||
|
||||
assert client is not None
|
||||
assert token_endpoint == "https://example.com/token"
|
||||
assert client.client_id == "test_client_id"
|
||||
assert client.client_secret == "test_client_secret"
|
||||
|
||||
def test_create_oauth2_session_oauth2_scheme(self):
|
||||
"""Test create_oauth2_session with OAuth2 scheme."""
|
||||
flows = OAuthFlows(
|
||||
authorizationCode=OAuthFlowAuthorizationCode(
|
||||
authorizationUrl="https://example.com/auth",
|
||||
tokenUrl="https://example.com/token",
|
||||
scopes={"read": "Read access", "write": "Write access"},
|
||||
)
|
||||
)
|
||||
scheme = OAuth2(type_="oauth2", flows=flows)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
redirect_uri="https://example.com/callback",
|
||||
),
|
||||
)
|
||||
|
||||
client, token_endpoint = create_oauth2_session(scheme, credential)
|
||||
|
||||
assert client is not None
|
||||
assert token_endpoint == "https://example.com/token"
|
||||
|
||||
def test_create_oauth2_session_invalid_scheme(self):
|
||||
"""Test create_oauth2_session with invalid scheme."""
|
||||
scheme = Mock() # Invalid scheme type
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OAUTH2,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
),
|
||||
)
|
||||
|
||||
client, token_endpoint = create_oauth2_session(scheme, credential)
|
||||
|
||||
assert client is None
|
||||
assert token_endpoint is None
|
||||
|
||||
def test_create_oauth2_session_missing_credentials(self):
|
||||
"""Test create_oauth2_session with missing credentials."""
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
type_="openIdConnect",
|
||||
openId_connect_url=(
|
||||
"https://example.com/.well-known/openid_configuration"
|
||||
),
|
||||
authorization_endpoint="https://example.com/auth",
|
||||
token_endpoint="https://example.com/token",
|
||||
scopes=["openid"],
|
||||
)
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
# Missing client_secret
|
||||
),
|
||||
)
|
||||
|
||||
client, token_endpoint = create_oauth2_session(scheme, credential)
|
||||
|
||||
assert client is None
|
||||
assert token_endpoint is None
|
||||
|
||||
def test_update_credential_with_tokens(self):
|
||||
"""Test update_credential_with_tokens function."""
|
||||
credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.OPEN_ID_CONNECT,
|
||||
oauth2=OAuth2Auth(
|
||||
client_id="test_client_id",
|
||||
client_secret="test_client_secret",
|
||||
),
|
||||
)
|
||||
|
||||
tokens = OAuth2Token({
|
||||
"access_token": "new_access_token",
|
||||
"refresh_token": "new_refresh_token",
|
||||
"expires_at": int(time.time()) + 3600,
|
||||
"expires_in": 3600,
|
||||
})
|
||||
|
||||
update_credential_with_tokens(credential, tokens)
|
||||
|
||||
assert credential.oauth2.access_token == "new_access_token"
|
||||
assert credential.oauth2.refresh_token == "new_refresh_token"
|
||||
assert credential.oauth2.expires_at == int(time.time()) + 3600
|
||||
assert credential.oauth2.expires_in == 3600
|
||||
+22
-18
@@ -20,6 +20,7 @@ from google.adk.auth.auth_credential import HttpAuth
|
||||
from google.adk.auth.auth_credential import HttpCredentials
|
||||
from google.adk.tools.application_integration_tool.integration_connector_tool import IntegrationConnectorTool
|
||||
from google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool import RestApiTool
|
||||
from google.adk.tools.openapi_tool.openapi_spec_parser.tool_auth_handler import AuthPreparationResult
|
||||
from google.genai.types import FunctionDeclaration
|
||||
from google.genai.types import Schema
|
||||
from google.genai.types import Type
|
||||
@@ -50,7 +51,9 @@ def mock_rest_api_tool():
|
||||
"required": ["user_id", "page_size", "filter", "connection_name"],
|
||||
}
|
||||
mock_tool._operation_parser = mock_parser
|
||||
mock_tool.call.return_value = {"status": "success", "data": "mock_data"}
|
||||
mock_tool.call = mock.AsyncMock(
|
||||
return_value={"status": "success", "data": "mock_data"}
|
||||
)
|
||||
return mock_tool
|
||||
|
||||
|
||||
@@ -179,9 +182,6 @@ async def test_run_with_auth_async_none_token(
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.ToolAuthHandler.from_tool_context"
|
||||
) as mock_from_tool_context:
|
||||
mock_tool_auth_handler_instance = mock.MagicMock()
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials.return_value.state = (
|
||||
"done"
|
||||
)
|
||||
# Simulate an AuthCredential that would cause _prepare_dynamic_euc to return None
|
||||
mock_auth_credential_without_token = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.HTTP,
|
||||
@@ -190,8 +190,12 @@ async def test_run_with_auth_async_none_token(
|
||||
credentials=HttpCredentials(token=None), # Token is None
|
||||
),
|
||||
)
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials.return_value.auth_credential = (
|
||||
mock_auth_credential_without_token
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials = mock.AsyncMock(
|
||||
return_value=(
|
||||
AuthPreparationResult(
|
||||
state="done", auth_credential=mock_auth_credential_without_token
|
||||
)
|
||||
)
|
||||
)
|
||||
mock_from_tool_context.return_value = mock_tool_auth_handler_instance
|
||||
|
||||
@@ -229,18 +233,18 @@ async def test_run_with_auth_async(
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.ToolAuthHandler.from_tool_context"
|
||||
) as mock_from_tool_context:
|
||||
mock_tool_auth_handler_instance = mock.MagicMock()
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials.return_value.state = (
|
||||
"done"
|
||||
)
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials.return_value.state = (
|
||||
"done"
|
||||
)
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials.return_value.auth_credential = AuthCredential(
|
||||
auth_type=AuthCredentialTypes.HTTP,
|
||||
http=HttpAuth(
|
||||
scheme="bearer",
|
||||
credentials=HttpCredentials(token="mocked_token"),
|
||||
),
|
||||
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials = mock.AsyncMock(
|
||||
return_value=AuthPreparationResult(
|
||||
state="done",
|
||||
auth_credential=AuthCredential(
|
||||
auth_type=AuthCredentialTypes.HTTP,
|
||||
http=HttpAuth(
|
||||
scheme="bearer",
|
||||
credentials=HttpCredentials(token="mocked_token"),
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
mock_from_tool_context.return_value = mock_tool_auth_handler_instance
|
||||
result = await integration_tool_with_auth.run_async(
|
||||
|
||||
@@ -14,6 +14,7 @@
|
||||
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import patch
|
||||
|
||||
@@ -194,7 +195,8 @@ class TestRestApiTool:
|
||||
@patch(
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.requests.request"
|
||||
)
|
||||
def test_call_success(
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_success(
|
||||
self,
|
||||
mock_request,
|
||||
mock_tool_context,
|
||||
@@ -217,7 +219,7 @@ class TestRestApiTool:
|
||||
)
|
||||
|
||||
# Call the method
|
||||
result = tool.call(args={}, tool_context=mock_tool_context)
|
||||
result = await tool.call(args={}, tool_context=mock_tool_context)
|
||||
|
||||
# Check the result
|
||||
assert result == {"result": "success"}
|
||||
@@ -225,7 +227,8 @@ class TestRestApiTool:
|
||||
@patch(
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.requests.request"
|
||||
)
|
||||
def test_call_auth_pending(
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_auth_pending(
|
||||
self,
|
||||
mock_request,
|
||||
sample_endpoint,
|
||||
@@ -246,12 +249,14 @@ class TestRestApiTool:
|
||||
"google.adk.tools.openapi_tool.openapi_spec_parser.rest_api_tool.ToolAuthHandler.from_tool_context"
|
||||
) as mock_from_tool_context:
|
||||
mock_tool_auth_handler_instance = MagicMock()
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials.return_value.state = (
|
||||
"pending"
|
||||
mock_prepare_result = MagicMock()
|
||||
mock_prepare_result.state = "pending"
|
||||
mock_tool_auth_handler_instance.prepare_auth_credentials = AsyncMock(
|
||||
return_value=mock_prepare_result
|
||||
)
|
||||
mock_from_tool_context.return_value = mock_tool_auth_handler_instance
|
||||
|
||||
response = tool.call(args={}, tool_context=None)
|
||||
response = await tool.call(args={}, tool_context=None)
|
||||
assert response == {
|
||||
"pending": True,
|
||||
"message": "Needs your authorization to access your data.",
|
||||
|
||||
@@ -116,7 +116,8 @@ def openid_connect_credential():
|
||||
return credential
|
||||
|
||||
|
||||
def test_openid_connect_no_auth_response(
|
||||
@pytest.mark.asyncio
|
||||
async def test_openid_connect_no_auth_response(
|
||||
openid_connect_scheme, openid_connect_credential
|
||||
):
|
||||
# Setup Mock exchanger
|
||||
@@ -132,12 +133,13 @@ def test_openid_connect_no_auth_response(
|
||||
credential_exchanger=mock_exchanger,
|
||||
credential_store=credential_store,
|
||||
)
|
||||
result = handler.prepare_auth_credentials()
|
||||
result = await handler.prepare_auth_credentials()
|
||||
assert result.state == 'pending'
|
||||
assert result.auth_credential == openid_connect_credential
|
||||
|
||||
|
||||
def test_openid_connect_with_auth_response(
|
||||
@pytest.mark.asyncio
|
||||
async def test_openid_connect_with_auth_response(
|
||||
openid_connect_scheme, openid_connect_credential, monkeypatch
|
||||
):
|
||||
mock_exchanger = MockOpenIdConnectCredentialExchanger(
|
||||
@@ -166,7 +168,7 @@ def test_openid_connect_with_auth_response(
|
||||
credential_exchanger=mock_exchanger,
|
||||
credential_store=credential_store,
|
||||
)
|
||||
result = handler.prepare_auth_credentials()
|
||||
result = await handler.prepare_auth_credentials()
|
||||
assert result.state == 'done'
|
||||
assert result.auth_credential.auth_type == AuthCredentialTypes.HTTP
|
||||
assert 'test_access_token' in result.auth_credential.http.credentials.token
|
||||
@@ -178,7 +180,8 @@ def test_openid_connect_with_auth_response(
|
||||
mock_auth_handler.get_auth_response.assert_called_once()
|
||||
|
||||
|
||||
def test_openid_connect_existing_token(
|
||||
@pytest.mark.asyncio
|
||||
async def test_openid_connect_existing_token(
|
||||
openid_connect_scheme, openid_connect_credential
|
||||
):
|
||||
_, existing_credential = token_to_scheme_credential(
|
||||
@@ -198,16 +201,17 @@ def test_openid_connect_existing_token(
|
||||
openid_connect_credential,
|
||||
credential_store=credential_store,
|
||||
)
|
||||
result = handler.prepare_auth_credentials()
|
||||
result = await handler.prepare_auth_credentials()
|
||||
assert result.state == 'done'
|
||||
assert result.auth_credential == existing_credential
|
||||
|
||||
|
||||
@patch(
|
||||
'google.adk.tools.openapi_tool.openapi_spec_parser.tool_auth_handler.OAuth2CredentialFetcher'
|
||||
'google.adk.tools.openapi_tool.openapi_spec_parser.tool_auth_handler.OAuth2CredentialRefresher'
|
||||
)
|
||||
def test_openid_connect_existing_oauth2_token_refresh(
|
||||
mock_oauth2_fetcher, openid_connect_scheme, openid_connect_credential
|
||||
@pytest.mark.asyncio
|
||||
async def test_openid_connect_existing_oauth2_token_refresh(
|
||||
mock_oauth2_refresher, openid_connect_scheme, openid_connect_credential
|
||||
):
|
||||
"""Test that OAuth2 tokens are refreshed when existing credentials are found."""
|
||||
# Create existing OAuth2 credential
|
||||
@@ -232,10 +236,13 @@ def test_openid_connect_existing_oauth2_token_refresh(
|
||||
),
|
||||
)
|
||||
|
||||
# Setup mock OAuth2CredentialFetcher
|
||||
mock_fetcher_instance = MagicMock()
|
||||
mock_fetcher_instance.refresh.return_value = refreshed_credential
|
||||
mock_oauth2_fetcher.return_value = mock_fetcher_instance
|
||||
# Setup mock OAuth2CredentialRefresher
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
mock_refresher_instance = MagicMock()
|
||||
mock_refresher_instance.is_refresh_needed = AsyncMock(return_value=True)
|
||||
mock_refresher_instance.refresh = AsyncMock(return_value=refreshed_credential)
|
||||
mock_oauth2_refresher.return_value = mock_refresher_instance
|
||||
|
||||
tool_context = create_mock_tool_context()
|
||||
credential_store = ToolContextCredentialStore(tool_context=tool_context)
|
||||
@@ -253,13 +260,17 @@ def test_openid_connect_existing_oauth2_token_refresh(
|
||||
credential_store=credential_store,
|
||||
)
|
||||
|
||||
result = handler.prepare_auth_credentials()
|
||||
result = await handler.prepare_auth_credentials()
|
||||
|
||||
# Verify OAuth2CredentialFetcher was called for refresh
|
||||
mock_oauth2_fetcher.assert_called_once_with(
|
||||
openid_connect_scheme, existing_credential
|
||||
# Verify OAuth2CredentialRefresher was called for refresh
|
||||
mock_oauth2_refresher.assert_called_once()
|
||||
|
||||
mock_refresher_instance.is_refresh_needed.assert_called_once_with(
|
||||
existing_credential
|
||||
)
|
||||
mock_refresher_instance.refresh.assert_called_once_with(
|
||||
existing_credential, openid_connect_scheme
|
||||
)
|
||||
mock_fetcher_instance.refresh.assert_called_once()
|
||||
|
||||
assert result.state == 'done'
|
||||
# The result should contain the refreshed credential after exchange
|
||||
|
||||
Reference in New Issue
Block a user