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
-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