mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
fix: oauth refresh not triggered on token expiry
Merge https://github.com/google/adk-python/pull/3767 Co-authored-by: Xuan Yang <xygoogle@google.com> COPYBARA_INTEGRATE_REVIEW=https://github.com/google/adk-python/pull/3767 from MarlzRana:marlzrana/fix-oauth-refresh-not-triggered-on-token-expiry 2dae3917de6e8d3fa5317857d06305d47c0773b0 PiperOrigin-RevId: 843756363
This commit is contained in:
committed by
Copybara-Service
parent
5c4bae7ff2
commit
69997cd5ef
@@ -33,7 +33,6 @@ 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(
|
||||
@@ -55,14 +54,14 @@ class TestOAuth2CredentialExchanger:
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
exchange_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"
|
||||
assert exchange_result.credential == credential
|
||||
assert exchange_result.credential.oauth2.access_token == "existing_token"
|
||||
assert not exchange_result.was_exchanged
|
||||
|
||||
@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
|
||||
@@ -96,14 +95,16 @@ class TestOAuth2CredentialExchanger:
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
exchange_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"
|
||||
assert exchange_result.credential.oauth2.access_token == "new_access_token"
|
||||
assert (
|
||||
exchange_result.credential.oauth2.refresh_token == "new_refresh_token"
|
||||
)
|
||||
assert exchange_result.was_exchanged
|
||||
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(
|
||||
@@ -122,7 +123,6 @@ class TestOAuth2CredentialExchanger:
|
||||
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
|
||||
@@ -146,14 +146,14 @@ class TestOAuth2CredentialExchanger:
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
exchange_result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Should return original credential when session creation fails
|
||||
assert result == credential
|
||||
assert result.oauth2.access_token is None
|
||||
assert exchange_result.credential == credential
|
||||
assert exchange_result.credential.oauth2.access_token is None
|
||||
assert not exchange_result.was_exchanged
|
||||
|
||||
@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
|
||||
@@ -181,14 +181,14 @@ class TestOAuth2CredentialExchanger:
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
exchange_result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Should return original credential when fetch_token fails
|
||||
assert result == credential
|
||||
assert result.oauth2.access_token is None
|
||||
assert exchange_result.credential == credential
|
||||
assert exchange_result.credential.oauth2.access_token is None
|
||||
assert not exchange_result.was_exchanged
|
||||
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(
|
||||
@@ -217,14 +217,14 @@ class TestOAuth2CredentialExchanger:
|
||||
"google.adk.auth.exchanger.oauth2_credential_exchanger.AUTHLIB_AVAILABLE",
|
||||
False,
|
||||
):
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
exchange_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
|
||||
assert exchange_result.credential == credential
|
||||
assert exchange_result.credential.oauth2.access_token is None
|
||||
assert not exchange_result.was_exchanged
|
||||
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Session")
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_client_credentials_success(self, mock_oauth2_session):
|
||||
"""Test successful client credentials exchange."""
|
||||
# Setup mock
|
||||
@@ -255,17 +255,19 @@ class TestOAuth2CredentialExchanger:
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
exchange_result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Verify client credentials exchange was successful
|
||||
assert result.oauth2.access_token == "client_access_token"
|
||||
assert (
|
||||
exchange_result.credential.oauth2.access_token == "client_access_token"
|
||||
)
|
||||
assert exchange_result.was_exchanged
|
||||
mock_client.fetch_token.assert_called_once_with(
|
||||
"https://example.com/token",
|
||||
grant_type="client_credentials",
|
||||
)
|
||||
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Session")
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_client_credentials_failure(self, mock_oauth2_session):
|
||||
"""Test client credentials exchange failure."""
|
||||
# Setup mock to raise exception during fetch_token
|
||||
@@ -292,15 +294,15 @@ class TestOAuth2CredentialExchanger:
|
||||
)
|
||||
|
||||
exchanger = OAuth2CredentialExchanger()
|
||||
result = await exchanger.exchange(credential, scheme)
|
||||
exchange_result = await exchanger.exchange(credential, scheme)
|
||||
|
||||
# Should return original credential when client credentials exchange fails
|
||||
assert result == credential
|
||||
assert result.oauth2.access_token is None
|
||||
assert exchange_result.credential == credential
|
||||
assert exchange_result.credential.oauth2.access_token is None
|
||||
assert not exchange_result.was_exchanged
|
||||
mock_client.fetch_token.assert_called_once()
|
||||
|
||||
@patch("google.adk.auth.oauth2_credential_util.OAuth2Session")
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchange_normalize_uri(self, mock_oauth2_session):
|
||||
"""Test exchange method normalizes auth_response_uri."""
|
||||
mock_client = Mock()
|
||||
@@ -344,7 +346,6 @@ class TestOAuth2CredentialExchanger:
|
||||
client_id="test_client_id",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_determine_grant_type_client_credentials(self):
|
||||
"""Test grant type determination for client credentials."""
|
||||
flows = OAuthFlows(
|
||||
@@ -361,7 +362,6 @@ class TestOAuth2CredentialExchanger:
|
||||
|
||||
assert grant_type == OAuthGrantType.CLIENT_CREDENTIALS
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_determine_grant_type_openid_connect(self):
|
||||
"""Test grant type determination for OpenID Connect (defaults to auth code)."""
|
||||
scheme = OpenIdConnectWithConfig(
|
||||
|
||||
Reference in New Issue
Block a user