Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions src/google/adk/flows/llm_flows/context/_contents.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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.
Expand All @@ -330,13 +333,22 @@ 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.
"""
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)
Expand Down Expand Up @@ -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,
Expand All @@ -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).
Expand All @@ -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)
Expand Down Expand Up @@ -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]:
Expand Down Expand Up @@ -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)
Expand All @@ -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,
Expand Down
25 changes: 18 additions & 7 deletions src/google/adk/workflow/_llm_agent_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -341,18 +341,22 @@ 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
branch = ctx._invocation_context.branch
if branch:
user_event.branch = branch
ctx.session.events.append(user_event)
return user_event


def process_llm_agent_output(
Expand Down Expand Up @@ -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}
Expand All @@ -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':
Expand Down
155 changes: 155 additions & 0 deletions tests/unittests/workflow/test_llm_agent_as_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Loading