fix: support fallback OAuth token and prefixless credential lookups

Merge https://github.com/google/adk-python/pull/5899

Resolve raw token strings and check for prefixless keys when retrieving authentication responses from session state.

Fixes #4712

PiperOrigin-RevId: 957435554
This commit is contained in:
Tony Coconate
2026-07-31 18:47:55 -07:00
committed by Copybara-Service
parent d9c5a129d8
commit f4e7233469
2 changed files with 191 additions and 6 deletions
+72 -6
View File
@@ -66,10 +66,13 @@ class AuthHandler:
return exchange_result.credential
async def parse_and_store_auth_response(self, state: State) -> None:
credential_key = self.auth_config.credential_key
if not credential_key:
raise ValueError("credential_key is empty.")
credential_key = "temp:" + self.auth_config.credential_key
temp_credential_key = "temp:" + credential_key
state[credential_key] = self.auth_config.exchanged_auth_credential
state[temp_credential_key] = self.auth_config.exchanged_auth_credential
if not isinstance(
self.auth_config.auth_scheme, SecurityBase
) or self.auth_config.auth_scheme.type_ not in (
@@ -78,15 +81,78 @@ class AuthHandler:
):
return
state[credential_key] = await self.exchange_auth_token()
state[temp_credential_key] = await self.exchange_auth_token()
def _validate(self) -> None:
if not self.auth_config.auth_scheme:
raise ValueError("auth_scheme is empty.")
def get_auth_response(self, state: State) -> AuthCredential:
credential_key = "temp:" + self.auth_config.credential_key
return state.get(credential_key, None)
def get_auth_response(self, state: State) -> AuthCredential | None:
# 1. Try reading the temp credential key (standard ADK flow)
credential_key = self.auth_config.credential_key
if not credential_key:
return None
temp_credential_key = "temp:" + credential_key
val = state.get(temp_credential_key, None)
if val is not None:
if isinstance(val, AuthCredential):
return val
if isinstance(val, dict):
return AuthCredential.model_validate(val)
if isinstance(val, str) and val:
return self._build_credential_from_string(val)
# 2. Try reading the credential key without the 'temp:' prefix
val = state.get(credential_key, None)
if val is not None:
if isinstance(val, AuthCredential):
return val
if isinstance(val, dict):
return AuthCredential.model_validate(val)
if isinstance(val, str) and val:
return self._build_credential_from_string(val)
return None
def _build_credential_from_string(self, val: str) -> AuthCredential:
from .auth_credential import AuthCredentialTypes
from .auth_credential import HttpAuth
from .auth_credential import HttpCredentials
from .auth_credential import OAuth2Auth
auth_scheme = self.auth_config.auth_scheme
if not auth_scheme:
return AuthCredential(
auth_type=AuthCredentialTypes.OAUTH2,
oauth2=OAuth2Auth(access_token=val),
)
scheme_type = auth_scheme.type_
if scheme_type == AuthSchemeType.apiKey:
return AuthCredential(
auth_type=AuthCredentialTypes.API_KEY,
api_key=val,
)
elif scheme_type == AuthSchemeType.http:
scheme = getattr(auth_scheme, "scheme", "bearer")
return AuthCredential(
auth_type=AuthCredentialTypes.HTTP,
http=HttpAuth(
scheme=scheme,
credentials=HttpCredentials(token=val),
),
)
elif scheme_type in (AuthSchemeType.oauth2, AuthSchemeType.openIdConnect):
return AuthCredential(
auth_type=AuthCredentialTypes.OAUTH2,
oauth2=OAuth2Auth(access_token=val),
)
else:
return AuthCredential(
auth_type=AuthCredentialTypes.OAUTH2,
oauth2=OAuth2Auth(access_token=val),
)
def generate_auth_request(self) -> AuthConfig:
if not isinstance(
+119
View File
@@ -607,6 +607,112 @@ class TestGetAuthResponse:
result = handler.get_auth_response(state)
assert result is None
def test_get_auth_response_temp_prefix_str_token(self, auth_config):
"""Test retrieving a string token stored under temp prefix in state."""
handler = AuthHandler(auth_config)
state = MockState()
credential_key = auth_config.credential_key
state["temp:" + credential_key] = "ya29.mock_token"
result = handler.get_auth_response(state)
assert result is not None
assert result.auth_type == AuthCredentialTypes.OAUTH2
assert result.oauth2.access_token == "ya29.mock_token"
def test_get_auth_response_no_prefix_credential(
self, auth_config, oauth2_credentials_with_auth_uri
):
"""Test retrieving a credential stored under the key without prefix."""
handler = AuthHandler(auth_config)
state = MockState()
credential_key = auth_config.credential_key
state[credential_key] = oauth2_credentials_with_auth_uri
result = handler.get_auth_response(state)
assert result == oauth2_credentials_with_auth_uri
def test_get_auth_response_no_prefix_str_token(self, auth_config):
"""Test retrieving a string token stored under the key without prefix."""
handler = AuthHandler(auth_config)
state = MockState()
credential_key = auth_config.credential_key
state[credential_key] = "ya29.mock_token_no_prefix"
result = handler.get_auth_response(state)
assert result is not None
assert result.auth_type == AuthCredentialTypes.OAUTH2
assert result.oauth2.access_token == "ya29.mock_token_no_prefix"
def test_get_auth_response_temp_prefix_dict(self, auth_config):
"""Test retrieving a credential dictionary stored under temp prefix."""
handler = AuthHandler(auth_config)
state = MockState()
credential_key = auth_config.credential_key
# Store dict in state representing an AuthCredential
state["temp:" + credential_key] = {
"auth_type": "oauth2",
"oauth2": {"access_token": "ya29.mock_token_from_dict"},
}
result = handler.get_auth_response(state)
assert result is not None
assert result.auth_type == AuthCredentialTypes.OAUTH2
assert result.oauth2.access_token == "ya29.mock_token_from_dict"
def test_get_auth_response_no_prefix_dict(self, auth_config):
"""Test retrieving a credential dictionary stored under the key without prefix."""
handler = AuthHandler(auth_config)
state = MockState()
credential_key = auth_config.credential_key
state[credential_key] = {
"auth_type": "oauth2",
"oauth2": {"access_token": "ya29.mock_token_from_dict_no_prefix"},
}
result = handler.get_auth_response(state)
assert result is not None
assert result.auth_type == AuthCredentialTypes.OAUTH2
assert result.oauth2.access_token == "ya29.mock_token_from_dict_no_prefix"
def test_get_auth_response_api_key_str(self):
"""Test retrieving a string token under apiKey scheme wraps it as APIKey."""
auth_scheme = APIKey(**{"name": "X-API-Key", "in": APIKeyIn.header})
config = AuthConfig(auth_scheme=auth_scheme)
handler = AuthHandler(config)
state = MockState()
credential_key = config.credential_key
state["temp:" + credential_key] = "my_api_key_value"
result = handler.get_auth_response(state)
assert result is not None
assert result.auth_type == AuthCredentialTypes.API_KEY
assert result.api_key == "my_api_key_value"
def test_get_auth_response_http_str(self):
"""Test retrieving a string token under http bearer scheme wraps it as HTTP Bearer."""
from fastapi.openapi.models import HTTPBearer
auth_scheme = HTTPBearer()
config = AuthConfig(auth_scheme=auth_scheme)
handler = AuthHandler(config)
state = MockState()
credential_key = config.credential_key
state["temp:" + credential_key] = "my_http_bearer_token"
result = handler.get_auth_response(state)
assert result is not None
assert result.auth_type == AuthCredentialTypes.HTTP
assert result.http is not None
assert result.http.scheme == "bearer"
assert result.http.credentials.token == "my_http_bearer_token"
class TestParseAndStoreAuthResponse:
"""Tests for the parse_and_store_auth_response method."""
@@ -650,6 +756,19 @@ class TestParseAndStoreAuthResponse:
assert state["temp:" + credential_key] == mock_exchange_token.return_value
assert mock_exchange_token.called
@pytest.mark.asyncio
async def test_empty_credential_key_raises_error(self, oauth2_auth_scheme):
"""Test that ValueError is raised when credential_key is empty."""
config = AuthConfig(
auth_scheme=oauth2_auth_scheme,
)
config.credential_key = "" # Bypass init logic that sets it
handler = AuthHandler(config)
state = MockState()
with pytest.raises(ValueError, match="credential_key is empty."):
await handler.parse_and_store_auth_response(state)
class TestExchangeAuthToken:
"""Tests for the exchange_auth_token method."""