From d0f7bda89db03f6bcc579cb092caf967b93155c9 Mon Sep 17 00:00:00 2001 From: Abhay Joshi Date: Sun, 27 Sep 2026 19:58:00 +0000 Subject: [PATCH] fix(workflow): isolate and clean up single_turn LlmAgent node_input events --- .../adk/flows/llm_flows/context/_contents.py | 18 ++ src/google/adk/workflow/_llm_agent_wrapper.py | 25 ++- .../workflow/test_llm_agent_as_node.py | 155 ++++++++++++++++++ 3 files changed, 191 insertions(+), 7 deletions(-) diff --git a/src/google/adk/flows/llm_flows/context/_contents.py b/src/google/adk/flows/llm_flows/context/_contents.py index cf6138ab58..d562178475 100644 --- a/src/google/adk/flows/llm_flows/context/_contents.py +++ b/src/google/adk/flows/llm_flows/context/_contents.py @@ -126,6 +126,7 @@ async def run_async( agent.name, preserve_function_call_ids=preserve_function_call_ids, isolation_scope=invocation_context.isolation_scope, + node_path=invocation_context.node_path, is_single_turn=is_single_turn, user_content=invocation_context.user_content, include_thoughts_from_other_agents=include_thoughts_from_other_agents, @@ -139,6 +140,7 @@ async def run_async( agent.name, preserve_function_call_ids=preserve_function_call_ids, isolation_scope=invocation_context.isolation_scope, + node_path=invocation_context.node_path, is_single_turn=is_single_turn, user_content=invocation_context.user_content, include_thoughts_from_other_agents=False, @@ -312,6 +314,7 @@ def _should_include_event_in_context( event: Event, isolation_scope: str | None = None, *, + node_path: str | None = None, include_thoughts: bool = False, ) -> bool: """Determines if an event should be included in the LLM context. @@ -330,6 +333,7 @@ def _should_include_event_in_context( current_branch: The current branch of the agent. event: The event to filter. isolation_scope: The agent's isolation_scope. None means unscoped. + node_path: The current workflow node path, if executing as a node. Returns: True if the event should be included in the context, False otherwise. @@ -337,6 +341,14 @@ def _should_include_event_in_context( ev_iso = getattr(event, 'isolation_scope', None) if ev_iso != isolation_scope: return False + ev_node_info = getattr(event, 'node_info', None) + ev_node_path = getattr(ev_node_info, 'path', None) if ev_node_info else None + if ( + event.author == 'user' + and ev_node_path + and ev_node_path != (node_path or '') + ): + return False return not ( _contains_empty_content(event, include_thoughts=include_thoughts) or not _is_event_belongs_to_branch(current_branch, event) @@ -399,6 +411,7 @@ def _get_contents( *, preserve_function_call_ids: bool = False, isolation_scope: str | None = None, + node_path: str | None = None, is_single_turn: bool = False, user_content: types.Content | None = None, include_thoughts_from_other_agents: bool = False, @@ -414,6 +427,7 @@ def _get_contents( preserve_function_call_ids: Whether to preserve function call ids. isolation_scope: scope tag — when set, restricts events to those with matching ``event.isolation_scope`` (or unscoped). + node_path: The current workflow node path, if executing as a node. user_content: Fallback first user turn for task agents whose originating delegation FC is not in session (workflow-node task case). @@ -440,6 +454,7 @@ def _get_contents( current_branch, e, isolation_scope=isolation_scope, + node_path=node_path, include_thoughts=( include_thoughts_from_other_agents and _is_other_agent_reply(agent_name, e) @@ -586,6 +601,7 @@ def _get_current_turn_contents( preserve_function_call_ids: bool = False, is_single_turn: bool = False, isolation_scope: str | None = None, + node_path: str | None = None, user_content: types.Content | None = None, include_thoughts_from_other_agents: bool = False, ) -> list[types.Content]: @@ -637,6 +653,7 @@ def _get_current_turn_contents( current_branch, event, isolation_scope=isolation_scope, + node_path=node_path, include_thoughts=( include_thoughts_from_other_agents and _is_other_agent_reply(agent_name, event) @@ -652,6 +669,7 @@ def _get_current_turn_contents( agent_name, preserve_function_call_ids=preserve_function_call_ids, isolation_scope=isolation_scope, + node_path=node_path, is_single_turn=is_single_turn, user_content=user_content, include_thoughts_from_other_agents=include_thoughts_from_other_agents, diff --git a/src/google/adk/workflow/_llm_agent_wrapper.py b/src/google/adk/workflow/_llm_agent_wrapper.py index 4818b7a1ff..13893e88b5 100644 --- a/src/google/adk/workflow/_llm_agent_wrapper.py +++ b/src/google/adk/workflow/_llm_agent_wrapper.py @@ -310,7 +310,7 @@ def prepare_llm_agent_context(agent: LlmAgent, ctx: Context) -> Context: def prepare_llm_agent_input( agent: LlmAgent, ctx: Context, node_input: object -) -> None: +) -> Event | None: """Prepares the input for running LlmAgent as a node. For ``single_turn`` mode, append a user-role event with the input @@ -341,11 +341,14 @@ def prepare_llm_agent_input( or agent.mode != 'single_turn' or bool(ctx.resume_inputs) ): - return + return None agent_input = to_user_content(node_input) user_event = Event(author='user', message=agent_input) if user_event.content is not None: user_event.content.role = 'user' + node_path = getattr(ctx, 'node_path', None) + if isinstance(node_path, str) and node_path: + user_event.node_info.path = node_path iso = getattr(ctx, 'isolation_scope', None) if iso: user_event.isolation_scope = iso @@ -353,6 +356,7 @@ def prepare_llm_agent_input( if branch: user_event.branch = branch ctx.session.events.append(user_event) + return user_event def process_llm_agent_output( @@ -411,7 +415,7 @@ async def run_llm_agent_as_node( agent.include_contents = 'none' agent_ctx = prepare_llm_agent_context(agent, ctx) - prepare_llm_agent_input(agent, agent_ctx, node_input) + injected_input_event = prepare_llm_agent_input(agent, agent_ctx, node_input) ic = agent_ctx.get_invocation_context() update: dict[str, object] = {'agent': agent} @@ -435,10 +439,17 @@ async def run_llm_agent_as_node( if agent.mode == 'single_turn': # is_live is always False here (single_turn forces non-live). - async with aclosing(agent.run_async(ic)) as run_iter: - async for event in run_iter: - process_llm_agent_output(agent, ctx, event) - yield event + try: + async with aclosing(agent.run_async(ic)) as run_iter: + async for event in run_iter: + process_llm_agent_output(agent, ctx, event) + yield event + finally: + if ( + injected_input_event is not None + and injected_input_event in agent_ctx.session.events + ): + agent_ctx.session.events.remove(injected_input_event) return if agent.mode == 'chat': diff --git a/tests/unittests/workflow/test_llm_agent_as_node.py b/tests/unittests/workflow/test_llm_agent_as_node.py index 15f70215a3..124359c83a 100644 --- a/tests/unittests/workflow/test_llm_agent_as_node.py +++ b/tests/unittests/workflow/test_llm_agent_as_node.py @@ -38,6 +38,7 @@ from google.adk.tools.function_tool import FunctionTool from google.adk.tools.long_running_tool import LongRunningFunctionTool from google.adk.workflow import _llm_agent_wrapper as agent_wrapper +from google.adk.workflow import node from google.adk.workflow import START from google.adk.workflow._llm_agent_wrapper import process_llm_agent_output from google.adk.workflow._workflow import Workflow @@ -1905,3 +1906,157 @@ def test_process_llm_agent_output_blank_schema_response_writes_no_state(): assert event.output is None assert ctx.actions.state_delta == {} + + +@pytest.mark.asyncio +async def test_single_turn_node_input_does_not_leak_across_sequential_tools( + request: pytest.FixtureRequest, +): + """Single-turn node_input must not leak into root agent across tool turns.""" + from . import testing_utils + + fake_pdf = b'%PDF-1.4-FAKE-BYTES' + worker_model = testing_utils.MockModel.create( + responses=['worker-summary-a', 'worker-summary-b'] + ) + worker = LlmAgent( + name='worker', + model=worker_model, + instruction='Summarize the attached document.', + mode='single_turn', + ) + + @node(name='run_worker', rerun_on_resume=True) + async def run_worker(ctx: Context, node_input: str) -> Any: + return await ctx.run_node( + worker, + node_input=types.Content( + role='user', + parts=[ + types.Part.from_text(text=f'INTERNAL-{node_input}'), + types.Part.from_bytes( + data=fake_pdf, mime_type='application/pdf' + ), + ], + ), + ) + + wf = Workflow( + name='doc_wf', + edges=[(START, run_worker)], + ) + + async def the_tool(label: str, tool_context: Context) -> dict[str, Any]: + out = await tool_context.run_node( + wf, node_input=label, run_id=f'run-{label}' + ) + return {'summary': f'done:{label}:{out}'} + + fc_a = types.Part.from_function_call(name='the_tool', args={'label': 'a'}) + fc_b = types.Part.from_function_call(name='the_tool', args={'label': 'b'}) + root_model = testing_utils.MockModel.create( + responses=[fc_a, fc_b, 'All tools completed.'] + ) + root_agent = LlmAgent( + name='root_agent', + model=root_model, + instruction='Call the_tool twice sequentially.', + tools=[the_tool], + ) + + runner = testing_utils.InMemoryRunner(root_agent) + await runner.run_async(testing_utils.get_user_content('run both tools')) + + # Worker received both inputs (text + inline PDF bytes). + assert len(worker_model.requests) == 2 + for expected_label, req in zip(['a', 'b'], worker_model.requests): + worker_texts = [ + p.text + for c in req.contents + for p in c.parts or [] + if p.text is not None + ] + worker_blobs = [ + p.inline_data.data + for c in req.contents + for p in c.parts or [] + if p.inline_data is not None + ] + assert any(f'INTERNAL-{expected_label}' in t for t in worker_texts) + assert fake_pdf in worker_blobs + + # Root agent made 3 LLM calls (initial -> after tool a -> after tool b). + # None of its requests should contain the worker's text or inline PDF. + assert len(root_model.requests) == 3 + for req in root_model.requests: + root_texts = [ + p.text + for c in req.contents + for p in c.parts or [] + if p.text is not None + ] + root_blobs = [ + p.inline_data + for c in req.contents + for p in c.parts or [] + if p.inline_data is not None + ] + assert not any('INTERNAL-' in t for t in root_texts) + assert not root_blobs + + +@pytest.mark.asyncio +async def test_parallel_single_turn_nodes_only_see_own_node_input( + request: pytest.FixtureRequest, +): + """Concurrent single_turn nodes sharing session.events see only own input.""" + import asyncio + + from . import testing_utils + + model_a = testing_utils.MockModel.create(responses=['out-a']) + model_b = testing_utils.MockModel.create(responses=['out-b']) + worker_a = LlmAgent( + name='worker_a', + model=model_a, + instruction='Worker A.', + mode='single_turn', + ) + worker_b = LlmAgent( + name='worker_b', + model=model_b, + instruction='Worker B.', + mode='single_turn', + ) + + @node(rerun_on_resume=True) + async def fanout(ctx: Context) -> dict[str, Any]: + res_a, res_b = await asyncio.gather( + ctx.run_node(worker_a, node_input='SECRET_FOR_A'), + ctx.run_node(worker_b, node_input='SECRET_FOR_B'), + ) + return {'a': res_a, 'b': res_b} + + wf = Workflow(name='parallel_wf', edges=[(START, fanout)]) + runner = _new_workflow_runner(wf, request.function.__name__) + await runner.run_async(testing_utils.get_user_content('start')) + + assert len(model_a.requests) == 1 + texts_a = [ + p.text + for c in model_a.requests[0].contents + for p in c.parts or [] + if p.text + ] + assert any('SECRET_FOR_A' in t for t in texts_a) + assert not any('SECRET_FOR_B' in t for t in texts_a) + + assert len(model_b.requests) == 1 + texts_b = [ + p.text + for c in model_b.requests[0].contents + for p in c.parts or [] + if p.text + ] + assert any('SECRET_FOR_B' in t for t in texts_b) + assert not any('SECRET_FOR_A' in t for t in texts_b)