Skip to content

Commit ce2e4ca

Browse files
google-genai-botcopybara-github
authored andcommitted
fix: Resolve scheduler leakage and make scheduler instantiation explicit
- Remove implicit on-demand fabrication of DynamicNodeScheduler during child Context derivation to avoid mutating parent context state during execution. - Always instantiate and set the DynamicNodeScheduler at the root level in the runner. - Clear `_workflow_scheduler` on execution exit (in both runner and workflow loops) to prevent scheduler lifetime leakage. - Update tests to explicitly set `rerun_on_resume = True` where resumption of mock nodes is expected. - Always rerun Workflow nodes on resume to allow their internal loops to correctly drive resumption of child nodes. - Use `weakref` for `_workflow_scheduler` in `Context` to prevent reference cycle between scheduler, tasks, and contexts. PiperOrigin-RevId: 945945253
1 parent da50578 commit ce2e4ca

6 files changed

Lines changed: 28 additions & 47 deletions

File tree

src/google/adk/agents/context.py

Lines changed: 8 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,6 @@
2020
from collections.abc import Sequence
2121
from typing import Any
2222
from typing import TYPE_CHECKING
23-
import weakref
2423

2524
from opentelemetry import context as context_api
2625
from typing_extensions import override
@@ -56,7 +55,13 @@ def _derive_scheduler(
5655
) -> ScheduleDynamicNode | None:
5756
"""Derives the dynamic node scheduler from the parent context."""
5857
if parent_ctx:
59-
return parent_ctx._workflow_scheduler
58+
scheduler = parent_ctx._workflow_scheduler
59+
if scheduler is None:
60+
from ..workflow._dynamic_node_scheduler import DynamicNodeScheduler
61+
from ..workflow._dynamic_node_scheduler import DynamicNodeState
62+
63+
scheduler = DynamicNodeScheduler(state=DynamicNodeState())
64+
return scheduler
6065
return None
6166

6267

@@ -117,8 +122,6 @@ class Context(ReadonlyContext):
117122
fields`` section are available.
118123
"""
119124

120-
_workflow_scheduler_ref: weakref.ref[ScheduleDynamicNode] | None = None
121-
122125
def __init__(
123126
self,
124127
invocation_context: InvocationContext,
@@ -200,10 +203,7 @@ def __init__(
200203
node=node,
201204
)
202205
self._resume_inputs = resume_inputs or {}
203-
scheduler = _derive_scheduler(parent_ctx)
204-
self._workflow_scheduler_ref = (
205-
weakref.ref(scheduler) if scheduler is not None else None
206-
)
206+
self._workflow_scheduler = _derive_scheduler(parent_ctx)
207207
self._node_rerun_on_resume = node.rerun_on_resume if node else True
208208
self._child_run_counters: dict[str, int] = {}
209209
self._attempt_count = attempt_count
@@ -264,19 +264,6 @@ def isolation_scope(self) -> str | None:
264264
def isolation_scope(self, value: str | None) -> None:
265265
self._isolation_scope = value
266266

267-
@property
268-
def _workflow_scheduler(self) -> ScheduleDynamicNode | None:
269-
"""The workflow scheduler associated with this context."""
270-
if self._workflow_scheduler_ref is not None:
271-
return self._workflow_scheduler_ref()
272-
return None
273-
274-
@_workflow_scheduler.setter
275-
def _workflow_scheduler(self, value: ScheduleDynamicNode | None) -> None:
276-
self._workflow_scheduler_ref = (
277-
weakref.ref(value) if value is not None else None
278-
)
279-
280267
@property
281268
def tool_confirmation(self) -> ToolConfirmation | None:
282269
"""The tool confirmation of the current tool call."""

src/google/adk/runners.py

Lines changed: 16 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -573,6 +573,11 @@ async def _run_node_async(
573573

574574
root_ctx = Context(ic)
575575
root_agent = node or self.agent
576+
is_agent = isinstance(self.agent, BaseAgent)
577+
has_sub_agents = is_agent and bool(
578+
getattr(self.agent, 'sub_agents', None)
579+
)
580+
use_scheduler = is_agent and has_sub_agents
576581

577582
# The root chat coordinator's isolation_scope stays None: its own
578583
# events (FCs, text, synthesized FRs from completed task
@@ -586,10 +591,11 @@ async def _run_node_async(
586591

587592
async def _drive_root_node():
588593
try:
589-
# Rehydration warning: DynamicNodeScheduler relies on session.events scanning.
590-
# Stateful live EUC/LRO streams may rehydrate freshly if not yet persisted.
591-
scheduler = DynamicNodeScheduler(state=_LoopState())
592-
root_ctx._workflow_scheduler = scheduler
594+
if use_scheduler:
595+
# Rehydration warning: DynamicNodeScheduler relies on session.events scanning.
596+
# Stateful live EUC/LRO streams may rehydrate freshly if not yet persisted.
597+
scheduler = DynamicNodeScheduler(state=_LoopState())
598+
root_ctx._workflow_scheduler = scheduler
593599

594600
try:
595601
await root_ctx._run_node_internal(
@@ -603,7 +609,6 @@ async def _drive_root_node():
603609
except DynamicNodeFailError as e:
604610
raise e.error
605611
finally:
606-
root_ctx._workflow_scheduler = None
607612
await ic._event_queue.put((done_sentinel, None))
608613

609614
task = asyncio.create_task(_drive_root_node())
@@ -674,6 +679,7 @@ async def _run_node_live(
674679
from .workflow._errors import DynamicNodeFailError
675680
from .workflow._errors import NodeInterruptedError
676681
from .workflow._workflow import _LoopState
682+
from .workflow._workflow import Workflow
677683

678684
ic = self._new_invocation_context_for_live(
679685
session,
@@ -684,12 +690,15 @@ async def _run_node_live(
684690

685691
root_ctx = Context(ic)
686692
root_agent = self.agent
693+
is_workflow = isinstance(root_agent, Workflow)
694+
687695
done_sentinel = object()
688696

689697
async def _drive_root_node():
690698
try:
691-
scheduler = DynamicNodeScheduler(state=_LoopState())
692-
root_ctx._workflow_scheduler = scheduler
699+
if is_workflow:
700+
scheduler = DynamicNodeScheduler(state=_LoopState())
701+
root_ctx._workflow_scheduler = scheduler
693702

694703
try:
695704
await root_ctx.run_node(
@@ -701,7 +710,6 @@ async def _drive_root_node():
701710
except DynamicNodeFailError as e:
702711
raise e.error
703712
finally:
704-
root_ctx._workflow_scheduler = None
705713
await ic._event_queue.put((done_sentinel, None))
706714

707715
task = asyncio.create_task(_drive_root_node())

src/google/adk/workflow/_workflow.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -245,7 +245,6 @@ async def _run_impl(
245245
try:
246246
await self._run_loop(loop_state, ctx)
247247
finally:
248-
ctx._workflow_scheduler = None
249248
await self._cleanup_all_tasks(loop_state)
250249

251250
if loop_state.error_shut_down:

src/google/adk/workflow/utils/_replay_interceptor.py

Lines changed: 1 addition & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -79,17 +79,10 @@ def check_interception(
7979
interrupts=set(current_run.state.interrupts),
8080
)
8181

82+
# Intercept executions based on historical session events (cross-turn replay).
8283
if not recovered:
8384
return InterceptionResult(should_run=True)
8485

85-
from .._workflow import Workflow
86-
87-
if isinstance(node, Workflow):
88-
return InterceptionResult(
89-
should_run=True,
90-
resume_inputs=recovered.resolved_responses,
91-
)
92-
9386
unresolved = recovered.interrupt_ids - recovered.resolved_ids
9487

9588
should_run = False

tests/unittests/agents/test_context.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -649,12 +649,13 @@ def test_derive_scheduler_with_parent_having_scheduler(self):
649649

650650
def test_derive_scheduler_with_parent_no_scheduler(self):
651651
from google.adk.agents.context import _derive_scheduler
652+
from google.adk.workflow._dynamic_node_scheduler import DynamicNodeScheduler
652653

653654
mock_parent = MagicMock()
654655
mock_parent._workflow_scheduler = None
655656

656657
scheduler = _derive_scheduler(mock_parent)
657-
assert scheduler is None
658+
assert isinstance(scheduler, DynamicNodeScheduler)
658659

659660

660661
class TestContextGetInvocationContext:

tests/unittests/runners/test_runner_node.py

Lines changed: 1 addition & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -456,7 +456,6 @@ async def test_standalone_node_resume():
456456
"""A standalone node resumes with resume_inputs from function response."""
457457

458458
class _Node(BaseNode):
459-
rerun_on_resume: bool = True
460459

461460
async def _run_impl(
462461
self, *, ctx: Context, node_input: Any
@@ -482,7 +481,6 @@ async def test_resume_preserves_original_user_content():
482481
"""On resume, Runner passes the original text as node_input, not the FR."""
483482

484483
class _Node(BaseNode):
485-
rerun_on_resume: bool = True
486484

487485
async def _run_impl(
488486
self, *, ctx: Context, node_input: Any
@@ -513,7 +511,6 @@ async def test_resume_populates_invocation_user_content():
513511
seen: list[Any] = []
514512

515513
class _Node(BaseNode):
516-
rerun_on_resume: bool = True
517514

518515
async def _run_impl(
519516
self, *, ctx: Context, node_input: Any
@@ -540,7 +537,6 @@ async def test_resume_by_invocation_id_populates_user_content():
540537
seen: list[Any] = []
541538

542539
class _Node(BaseNode):
543-
rerun_on_resume: bool = True
544540

545541
async def _run_impl(
546542
self, *, ctx: Context, node_input: Any
@@ -568,10 +564,7 @@ async def _run_impl(
568564
invocation_id = updated.events[0].invocation_id
569565

570566
async for _ in runner.run_async(
571-
user_id='u',
572-
session_id=session.id,
573-
invocation_id=invocation_id,
574-
new_message=_make_resume_message(fc_name='tool', response={'v': 1}),
567+
user_id='u', session_id=session.id, invocation_id=invocation_id
575568
):
576569
pass
577570

0 commit comments

Comments
 (0)