From 2277516844c3e6dbbfd3a7974495659303846325 Mon Sep 17 00:00:00 2001 From: George Weale Date: Mon, 17 Aug 2026 23:01:12 +0000 Subject: [PATCH 1/2] fix: use branch-scoped events for auth responses (v1) Port of the same fix on main. The auth preprocessor read the whole session event list when looking for the adk_request_credential call that a credential response belongs to. During a parallel or delegated execution, a credential returned for one branch matched the pending request on a sibling branch, and the preprocessor re-executed that branch's tool. It now reads invocation_context._get_events(current_branch=True), so only events on the current branch, plus the user's own events, are considered. The user's credential reply carries no branch and stays visible everywhere, so an ordinary same-branch resume is unaffected. Behaviour change: an adk_request_credential response no longer resumes a tool whose auth request was issued on a different branch. On v1 branches are compared for equality rather than by prefix, so this also covers a request issued on a branch below the current one; a test pins that boundary. --- src/google/adk/auth/auth_preprocessor.py | 2 +- .../unittests/auth/test_auth_preprocessor.py | 187 +++++++++++++++++- 2 files changed, 184 insertions(+), 5 deletions(-) diff --git a/src/google/adk/auth/auth_preprocessor.py b/src/google/adk/auth/auth_preprocessor.py index 76dd2ddab46..c5d709af7ed 100644 --- a/src/google/adk/auth/auth_preprocessor.py +++ b/src/google/adk/auth/auth_preprocessor.py @@ -134,7 +134,7 @@ async def run_async( agent = invocation_context.agent if not hasattr(agent, 'canonical_tools'): return - events = invocation_context.session.events + events = invocation_context._get_events(current_branch=True) if not events: return diff --git a/tests/unittests/auth/test_auth_preprocessor.py b/tests/unittests/auth/test_auth_preprocessor.py index fb45cc34ac8..175d124de1e 100644 --- a/tests/unittests/auth/test_auth_preprocessor.py +++ b/tests/unittests/auth/test_auth_preprocessor.py @@ -20,7 +20,12 @@ from unittest.mock import Mock from unittest.mock import patch +from fastapi.openapi.models import APIKey +from fastapi.openapi.models import APIKeyIn from google.adk.agents.invocation_context import InvocationContext +from google.adk.agents.llm_agent import LlmAgent +from google.adk.auth.auth_credential import AuthCredential +from google.adk.auth.auth_credential import AuthCredentialTypes from google.adk.auth.auth_handler import AuthHandler from google.adk.auth.auth_preprocessor import _AuthLlmRequestProcessor from google.adk.auth.auth_tool import AuthConfig @@ -28,6 +33,9 @@ from google.adk.events.event import Event from google.adk.flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME from google.adk.models.llm_request import LlmRequest +from google.adk.sessions.base_session_service import BaseSessionService +from google.adk.sessions.session import Session +from google.genai import types import pytest @@ -42,8 +50,6 @@ def processor(self): @pytest.fixture def mock_llm_agent(self): """Create a mock LlmAgent.""" - from google.adk.agents.llm_agent import LlmAgent - agent = Mock(spec=LlmAgent) agent.canonical_tools = AsyncMock(return_value=[]) return agent @@ -69,6 +75,7 @@ def mock_invocation_context(self, mock_llm_agent, mock_session): context = Mock(spec=InvocationContext) context.agent = mock_llm_agent context.session = mock_session + context._get_events.side_effect = lambda **_: context.session.events return context @pytest.fixture @@ -165,7 +172,8 @@ async def test_non_llm_agent_returns_early( ): """Test that non-LLM agents return early.""" mock_context = Mock(spec=InvocationContext) - mock_context.agent = Mock() + # Using spec=[] ensures hasattr(agent, 'canonical_tools') returns False. + mock_context.agent = Mock(spec=[]) mock_context.agent.__class__.__name__ = 'BaseAgent' mock_context.session = mock_session @@ -273,6 +281,38 @@ async def test_last_event_no_auth_responses_returns_early( assert result == [] + @pytest.mark.asyncio + @patch('google.adk.auth.auth_preprocessor.AuthHandler') + @patch('google.adk.auth.auth_tool.AuthConfig.model_validate') + async def test_ignores_auth_responses_outside_current_branch( + self, + mock_auth_config_validate, + mock_auth_handler_class, + processor, + mock_invocation_context, + mock_llm_request, + mock_user_event_with_auth_response, + ): + """Test auth responses hidden by branch filtering are ignored.""" + mock_invocation_context.session.events = [ + mock_user_event_with_auth_response + ] + mock_invocation_context._get_events.side_effect = None + mock_invocation_context._get_events.return_value = [] + + result = [] + async for event in processor.run_async( + mock_invocation_context, mock_llm_request + ): + result.append(event) + + mock_invocation_context._get_events.assert_called_once_with( + current_branch=True + ) + mock_auth_config_validate.assert_not_called() + mock_auth_handler_class.assert_not_called() + assert result == [] + @pytest.mark.asyncio @patch('google.adk.auth.auth_preprocessor.AuthHandler') @patch('google.adk.auth.auth_tool.AuthConfig.model_validate') @@ -536,7 +576,8 @@ async def test_isinstance_check_for_llm_agent( # Create a mock that fails isinstance check mock_context = Mock(spec=InvocationContext) - mock_context.agent = Mock() # This will fail isinstance(agent, LlmAgent) + # This will fail isinstance(agent, LlmAgent) + mock_context.agent = Mock(spec=[]) mock_context.session = mock_session result = [] @@ -544,3 +585,141 @@ async def test_isinstance_check_for_llm_agent( result.append(event) assert result == [] + + +class TestBranchScopedAuthResponses: + """Tests the branch filter itself, rather than a mocked ``_get_events``. + + The tests above build the invocation context with ``Mock``, so + ``_get_events`` is a mock and the predicate inside + ``InvocationContext._get_events`` never runs. These tests build a real + context and real events, because that predicate is what decides whether a + credential requested on one branch can resume a tool on another. + """ + + AUTH_FC_ID = 'auth_fc_id' + TOOL_FC_ID = 'tool_fc_id' + + @pytest.fixture + def auth_config(self): + """An API key config, which is stored without a token exchange.""" + return AuthConfig( + auth_scheme=APIKey(**{'name': 'test_api_key', 'in': APIKeyIn.header}), + raw_auth_credential=AuthCredential( + auth_type=AuthCredentialTypes.API_KEY, api_key='test_api_key' + ), + ) + + def _build_events(self, auth_config, request_branch): + """Builds a paused tool call, its credential request and the user reply.""" + tool_call = types.FunctionCall( + id=self.TOOL_FC_ID, name='some_tool', args={} + ) + tool_call_event = Event( + invocation_id='inv_1', + author='test_agent', + branch=request_branch, + content=types.Content( + role='model', parts=[types.Part(function_call=tool_call)] + ), + ) + auth_request = types.FunctionCall( + id=self.AUTH_FC_ID, + name=REQUEST_EUC_FUNCTION_CALL_NAME, + args=AuthToolArguments( + function_call_id=self.TOOL_FC_ID, auth_config=auth_config + ).model_dump(exclude_none=True, by_alias=True), + ) + auth_request_event = Event( + invocation_id='inv_1', + author='test_agent', + branch=request_branch, + content=types.Content( + role='model', parts=[types.Part(function_call=auth_request)] + ), + long_running_tool_ids={self.AUTH_FC_ID}, + ) + auth_response = types.FunctionResponse( + id=self.AUTH_FC_ID, + name=REQUEST_EUC_FUNCTION_CALL_NAME, + response=auth_config.model_dump(exclude_none=True, by_alias=True), + ) + # The client sends the credential back with no branch, so it stays visible + # from every branch. + auth_response_event = Event( + invocation_id='inv_1', + author='user', + content=types.Content( + role='user', parts=[types.Part(function_response=auth_response)] + ), + ) + return [tool_call_event, auth_request_event, auth_response_event] + + async def _run_processor(self, events, branch): + """Runs the processor and returns the mocked tool re-execution call.""" + invocation_context = InvocationContext( + session_service=Mock(spec=BaseSessionService), + agent=LlmAgent(name='test_agent'), + invocation_id='inv_1', + session=Mock(spec=Session, events=events, state={}), + branch=branch, + ) + with patch( + 'google.adk.flows.llm_flows.functions.handle_function_calls_async' + ) as mock_handle_function_calls: + mock_handle_function_calls.return_value = None + async for _ in _AuthLlmRequestProcessor().run_async( + invocation_context, Mock(spec=LlmRequest) + ): + pass + return mock_handle_function_calls + + @pytest.mark.asyncio + async def test_same_branch_auth_response_resumes_the_tool(self, auth_config): + """The ordinary case: the request and the resume share a branch.""" + events = self._build_events(auth_config, request_branch='root_agent') + + mock_handle_function_calls = await self._run_processor( + events, branch='root_agent' + ) + + mock_handle_function_calls.assert_called_once() + call_args = mock_handle_function_calls.call_args + assert call_args[0][1] is events[0] + assert call_args[0][3] == {self.TOOL_FC_ID} + + @pytest.mark.asyncio + async def test_sibling_branch_auth_response_does_not_resume_the_tool( + self, auth_config + ): + """A credential must not re-execute a tool paused on another branch.""" + events = self._build_events( + auth_config, request_branch='root_agent.branch_b' + ) + + mock_handle_function_calls = await self._run_processor( + events, branch='root_agent.branch_a' + ) + + mock_handle_function_calls.assert_not_called() + + @pytest.mark.asyncio + async def test_descendant_branch_auth_response_does_not_resume_the_tool( + self, auth_config + ): + """A request from a branch below the current one is also out of scope. + + Branches are compared for equality rather than by prefix, so a parent + branch does not resume a tool that a sub-agent paused. In practice both + strings come from the same invocation context, so they match; this pins + the boundary so that a change to the predicate is not silent. + """ + events = self._build_events( + auth_config, request_branch='root_agent.sub_agent' + ) + + mock_handle_function_calls = await self._run_processor( + events, branch='root_agent' + ) + + mock_handle_function_calls.assert_not_called() From 05337af7e17f02790fdb61bdd6365b5ec6fdfe5f Mon Sep 17 00:00:00 2001 From: George Weale Date: Mon, 17 Aug 2026 23:02:22 +0000 Subject: [PATCH 2/2] feat: forward the OAuth2 nonce to the authorization request (v1) Port of the same fix on main. OAuth2Auth already has a nonce field, described as the value that binds the user's session to the authorization request, but generate_auth_uri never put it in the query parameters it hands to create_authorization_url. A caller who set a nonce got an authorization request without one, so the returned id_token could not be bound to it and the OIDC replay defence was silently absent. The nonce is now passed through when the credential carries one, alongside the audience parameter that is handled the same way. Flows that leave it unset are unchanged. --- src/google/adk/auth/auth_handler.py | 2 + tests/unittests/auth/test_auth_handler.py | 53 +++++++++++++++++++++++ 2 files changed, 55 insertions(+) diff --git a/src/google/adk/auth/auth_handler.py b/src/google/adk/auth/auth_handler.py index 8e8f5d340b6..d01b0f5269b 100644 --- a/src/google/adk/auth/auth_handler.py +++ b/src/google/adk/auth/auth_handler.py @@ -201,6 +201,8 @@ def generate_auth_uri( } if auth_credential.oauth2.audience: params["audience"] = auth_credential.oauth2.audience + if auth_credential.oauth2.nonce: + params["nonce"] = auth_credential.oauth2.nonce # If using PKCE with S256, ensure a code_verifier exists. # If not provided in the credential, generate a cryptographically secure diff --git a/tests/unittests/auth/test_auth_handler.py b/tests/unittests/auth/test_auth_handler.py index c19a5d93fd3..2fa1228bf6d 100644 --- a/tests/unittests/auth/test_auth_handler.py +++ b/tests/unittests/auth/test_auth_handler.py @@ -304,6 +304,59 @@ def test_generate_auth_uri_pkce( assert "code_verifier" in kwargs assert kwargs["code_verifier"] == result.oauth2.code_verifier + @patch("google.adk.auth.auth_handler.OAuth2Session") + def test_generate_auth_uri_with_nonce( + self, mock_oauth2_session, oauth2_auth_scheme, oauth2_credentials + ): + """Test that a nonce is forwarded to the authorization request.""" + oauth2_credentials.oauth2.nonce = "test_nonce" + exchanged = oauth2_credentials.model_copy(deep=True) + + config = AuthConfig( + auth_scheme=oauth2_auth_scheme, + raw_auth_credential=oauth2_credentials, + exchanged_auth_credential=exchanged, + ) + + mock_client = Mock() + mock_oauth2_session.return_value = mock_client + mock_client.create_authorization_url.return_value = ( + "https://example.com/oauth2/authorize?nonce=test_nonce", + "mock_state", + ) + + handler = AuthHandler(config) + handler.generate_auth_uri() + + _, kwargs = mock_client.create_authorization_url.call_args + assert kwargs["nonce"] == "test_nonce" + + @patch("google.adk.auth.auth_handler.OAuth2Session") + def test_generate_auth_uri_without_nonce( + self, mock_oauth2_session, oauth2_auth_scheme, oauth2_credentials + ): + """Test that no nonce is sent when the credential has none.""" + exchanged = oauth2_credentials.model_copy(deep=True) + + config = AuthConfig( + auth_scheme=oauth2_auth_scheme, + raw_auth_credential=oauth2_credentials, + exchanged_auth_credential=exchanged, + ) + + mock_client = Mock() + mock_oauth2_session.return_value = mock_client + mock_client.create_authorization_url.return_value = ( + "https://example.com/oauth2/authorize", + "mock_state", + ) + + handler = AuthHandler(config) + handler.generate_auth_uri() + + _, kwargs = mock_client.create_authorization_url.call_args + assert "nonce" not in kwargs + def test_generate_auth_uri_unsupported_pkce_method( self, oauth2_auth_scheme, oauth2_credentials ):