diff --git a/src/google/adk/runners.py b/src/google/adk/runners.py index 626b6ee5..6d9f7007 100644 --- a/src/google/adk/runners.py +++ b/src/google/adk/runners.py @@ -652,109 +652,124 @@ class Runner: # on_run_error_callback: they are part of runner execution even though # they run before the main event loop. Notification-only; the original # exception is always re-raised, and after_run stays success-only. - try: - # Run callbacks on user message - if new_message: - modified_user_message = ( - await ic.plugin_manager.run_on_user_message_callback( - invocation_context=ic, user_message=new_message - ) - ) - if modified_user_message is not None: - new_message = modified_user_message - ic.user_content = new_message - - # Append user message to session for history - if new_message: - user_event = await self._append_user_event( - ic, new_message, state_delta=state_delta - ) - if yield_user_message and user_event: - yield user_event - - # Run before_run callbacks - await ic.plugin_manager.run_before_run_callback(invocation_context=ic) - except Exception as e: - await _notify_run_error(ic.plugin_manager, ic, e) - raise - - # 3. Start root node in background - from .agents.context import Context - from .workflow._dynamic_node_scheduler import DynamicNodeScheduler - from .workflow._errors import DynamicNodeFailError - from .workflow._errors import NodeInterruptedError - from .workflow._workflow import _LoopState - - root_ctx = Context(ic) - root_node = node or self.agent - is_agent = isinstance(self.agent, BaseAgent) - has_sub_agents = is_agent and bool( - getattr(self.agent, 'sub_agents', None) - ) - use_scheduler = is_agent and has_sub_agents - - # The root chat coordinator's isolation_scope stays None: its own - # events (FCs, text, synthesized FRs from completed task - # delegations) are also unscoped, so the content-builder's - # isolation_scope filter lets the coordinator see all of them - # across user turns. Task sub-agents are scoped under their - # originating function-call id and so remain invisible to the - # coordinator's view. - - done_sentinel = object() - - async def _drive_root_node() -> None: - try: - if use_scheduler: - # Rehydration warning: DynamicNodeScheduler relies on session.events scanning. - # Stateful live EUC/LRO streams may rehydrate freshly if not yet persisted. - scheduler = DynamicNodeScheduler(state=_LoopState()) - root_ctx._workflow_scheduler = scheduler - - try: - await root_ctx._run_node_internal( - root_node, - node_input=node_input, - resume_inputs=resume_inputs, - ) - except NodeInterruptedError: - # The node was interrupted (e.g. for HITL). - pass - except DynamicNodeFailError as e: - raise e.error - finally: - await ic._event_queue.put((done_sentinel, None)) - - task = asyncio.create_task(_drive_root_node()) - - # 4. Main loop: consume events, persist, yield run_error = None try: try: - async with aclosing( - self._consume_event_queue(ic, done_sentinel) - ) as agen: - async for event in agen: - yield event - finally: - # _cleanup_root_task re-raises a root-node Exception (if any) after - # the event stream has drained. - await self._cleanup_root_task(task, self.agent.name) - except Exception as e: - # An unhandled exception escaped runner execution. Notify plugins - # (notification-only) and re-raise. after_run stays success-only. - run_error = e - await _notify_run_error(ic.plugin_manager, ic, e) - raise + # Run callbacks on user message + if new_message: + modified_user_message = ( + await ic.plugin_manager.run_on_user_message_callback( + invocation_context=ic, user_message=new_message + ) + ) + if modified_user_message is not None: + new_message = modified_user_message + ic.user_content = new_message + + # Append user message to session for history + if new_message: + user_event = await self._append_user_event( + ic, new_message, state_delta=state_delta + ) + if yield_user_message and user_event: + yield user_event + + # Run before_run callbacks. A returned Content halts execution and ends + # the run with that content (same contract as the non-workflow path). + early_exit_result = await ic.plugin_manager.run_before_run_callback( + invocation_context=ic + ) + if isinstance(early_exit_result, types.Content): + early_exit_event = Event( + invocation_id=ic.invocation_id, + author='model', + content=early_exit_result, + ) + _apply_run_config_custom_metadata(early_exit_event, ic.run_config) + if self._should_append_event( + early_exit_event, is_live_call=False + ): + await self.session_service.append_event( + session=ic.session, + event=early_exit_event, + ) + yield early_exit_event + else: + # 3. Start root node in background + from .agents.context import Context + from .workflow._dynamic_node_scheduler import DynamicNodeScheduler + from .workflow._errors import DynamicNodeFailError + from .workflow._errors import NodeInterruptedError + from .workflow._workflow import _LoopState + + root_ctx = Context(ic) + root_node = node or self.agent + is_agent = isinstance(self.agent, BaseAgent) + has_sub_agents = is_agent and bool( + getattr(self.agent, 'sub_agents', None) + ) + use_scheduler = is_agent and has_sub_agents + + # The root chat coordinator's isolation_scope stays None: its own + # events (FCs, text, synthesized FRs from completed task + # delegations) are also unscoped, so the content-builder's + # isolation_scope filter lets the coordinator see all of them + # across user turns. Task sub-agents are scoped under their + # originating function-call id and so remain invisible to the + # coordinator's view. + + done_sentinel = object() + + async def _drive_root_node() -> None: + try: + if use_scheduler: + # Rehydration warning: DynamicNodeScheduler relies on session.events scanning. + # Stateful live EUC/LRO streams may rehydrate freshly if not yet persisted. + scheduler = DynamicNodeScheduler(state=_LoopState()) + root_ctx._workflow_scheduler = scheduler + + try: + await root_ctx._run_node_internal( + root_node, + node_input=node_input, + resume_inputs=resume_inputs, + ) + except NodeInterruptedError: + # The node was interrupted (e.g. for HITL). + pass + except DynamicNodeFailError as e: + raise e.error + finally: + await ic._event_queue.put((done_sentinel, None)) + + task = asyncio.create_task(_drive_root_node()) + + # 4. Main loop: consume events, persist, yield + try: + async with aclosing( + self._consume_event_queue(ic, done_sentinel) + ) as agen: + async for event in agen: + yield event + finally: + # _cleanup_root_task re-raises a root-node Exception (if any) after + # the event stream has drained. + await self._cleanup_root_task(task, self.agent.name) + except Exception as e: + # An unhandled exception escaped runner execution. Notify plugins + # (notification-only) and re-raise. after_run stays success-only. + run_error = e + await _notify_run_error(ic.plugin_manager, ic, e) + raise finally: # Success path (also caller early-stop via GeneratorExit, which is not # an Exception): run after_run and compaction. _cleanup_root_task has - # already run in the inner finally above. A failure in this success - # cleanup (e.g. an after_run plugin raising, which PluginManager - # surfaces as a RuntimeError) is itself an unhandled runner error, so - # notify on_run_error_callback once and re-raise. on_run_error is - # notification-only and never raises, so there is no recursive - # notification. + # already run in the inner finally above when a root task was created. + # A failure in this success cleanup (e.g. an after_run plugin raising, + # which PluginManager surfaces as a RuntimeError) is itself an + # unhandled runner error, so notify on_run_error_callback once and + # re-raise. on_run_error is notification-only and never raises, so + # there is no recursive notification. if run_error is None: try: await ic.plugin_manager.run_after_run_callback( diff --git a/tests/unittests/workflow/test_workflow_failures.py b/tests/unittests/workflow/test_workflow_failures.py index 8721acd8..4cbd63a2 100644 --- a/tests/unittests/workflow/test_workflow_failures.py +++ b/tests/unittests/workflow/test_workflow_failures.py @@ -21,8 +21,10 @@ from unittest import mock from google.adk import platform as adk_platform from google.adk.agents.context import Context +from google.adk.agents.invocation_context import InvocationContext from google.adk.apps.app import App from google.adk.events.event import Event +from google.adk.plugins.base_plugin import BasePlugin # Added for the moved test from google.adk.runners import Runner from google.adk.sessions.in_memory_session_service import InMemorySessionService @@ -1162,3 +1164,121 @@ async def test_multiple_failures_first_error_wins( with pytest.raises(ValueError, match='Fail 1'): await runner.run_async(testing_utils.get_user_content('start')) + + +@pytest.mark.asyncio +async def test_workflow_halts_when_before_run_callback_returns_content(): + """Regression for #6013: a plugin before_run_callback returning Content must + halt the workflow run with that content and skip node execution.""" + + ran = {'node': False} + + class _RecordingNode(BaseNode): + + @override + async def run( + self, *, ctx: Context, node_input: Any + ) -> AsyncGenerator[Any, None]: + ran['node'] = True + yield Event(output='should not run') + + class _HaltPlugin(BasePlugin): + + def __init__(self): + super().__init__(name='halt_plugin') + + async def before_run_callback( + self, *, invocation_context: InvocationContext + ) -> types.Content: + return types.Content( + role='model', parts=[types.Part(text='halted by plugin')] + ) + + graph = Graph(edges=[Edge(from_node=START, to_node=_RecordingNode(name='A'))]) + wf = Workflow(name='halt_wf', graph=graph) + + ss = InMemorySessionService() + app = App(name='test', root_agent=wf, plugins=[_HaltPlugin()]) + runner = Runner(app=app, session_service=ss) + session = await ss.create_session(app_name='test', user_id='u') + msg = types.Content(parts=[types.Part(text='start')], role='user') + events = [ + event + async for event in runner.run_async( + user_id='u', session_id=session.id, new_message=msg + ) + ] + + # The run halts with the plugin's content and the node never executes. + assert ran['node'] is False + assert any( + e.content + and e.content.parts + and e.content.parts[0].text == 'halted by plugin' + for e in events + ) + + +@pytest.mark.asyncio +async def test_workflow_dispatches_after_run_callback_on_before_run_early_exit( + monkeypatch: pytest.MonkeyPatch, +): + """A plugin after_run_callback and post-invocation compaction must be dispatched even when before_run_callback early exits with Content.""" + + ran = {'node': False, 'after_run': False, 'compaction': False} + + class _RecordingNode(BaseNode): + + @override + async def run( + self, *, ctx: Context, node_input: Any + ) -> AsyncGenerator[Any, None]: + ran['node'] = True + yield Event(output='should not run') + + class _HaltWithAfterRunPlugin(BasePlugin): + + def __init__(self): + super().__init__(name='halt_with_after_run_plugin') + + async def before_run_callback( + self, *, invocation_context: InvocationContext + ) -> types.Content: + return types.Content( + role='model', parts=[types.Part(text='halted by plugin')] + ) + + async def after_run_callback( + self, *, invocation_context: InvocationContext + ) -> None: + ran['after_run'] = True + + graph = Graph(edges=[Edge(from_node=START, to_node=_RecordingNode(name='A'))]) + wf = Workflow(name='halt_wf', graph=graph) + + ss = InMemorySessionService() + app = App(name='test', root_agent=wf, plugins=[_HaltWithAfterRunPlugin()]) + runner = Runner(app=app, session_service=ss) + + original_compaction = runner._run_post_invocation_compaction + + async def _mock_compaction(*args, **kwargs): + ran['compaction'] = True + await original_compaction(*args, **kwargs) + + monkeypatch.setattr( + runner, '_run_post_invocation_compaction', _mock_compaction + ) + + session = await ss.create_session(app_name='test', user_id='u') + msg = types.Content(parts=[types.Part(text='start')], role='user') + events = [ + event + async for event in runner.run_async( + user_id='u', session_id=session.id, new_message=msg + ) + ] + + assert ran['node'] is False + assert ran['after_run'] is True + assert ran['compaction'] is True