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:
Marlin Ranasinghe
2025-12-12 10:57:58 -08:00
committed by Copybara-Service
parent 5c4bae7ff2
commit 69997cd5ef
5 changed files with 75 additions and 59 deletions
@@ -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(