Skip to content

Commit 37aa730

Browse files
DeanChensjcopybara-github
authored andcommitted
fix(agents): trigger after_agent_callback on cancellation
When agent execution was cancelled via asyncio.CancelledError (e.g. client disconnect or RPC timeout), _handle_after_agent_callback was bypassed, causing execution metrics and timers in plugins to leak. This adds an explicit exception handler for CancelledError to ensure after_agent_callback runs for side-effects during cancellation, while preserving the invariant that end_invocation short-circuits and unhandled errors do not trigger after_agent_callback. Co-authored-by: Shangjie Chen <deanchen@google.com> PiperOrigin-RevId: 980466771
1 parent c03683a commit 37aa730

2 files changed

Lines changed: 392 additions & 2 deletions

File tree

‎src/google/adk/agents/base_agent.py‎

Lines changed: 48 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from __future__ import annotations
1616

1717
import abc
18+
import asyncio
1819
import inspect
1920
import logging
2021
from typing import Any
@@ -341,8 +342,12 @@ async def run_async(
341342
async def _run() -> AsyncGenerator[Event, None]:
342343
ctx = self._create_invocation_context(parent_context)
343344
async with _instrumentation.record_agent_invocation(ctx, self):
345+
before_callback_completed = False
346+
after_callback_called = False
344347
try:
345-
if event := await self._handle_before_agent_callback(ctx):
348+
event = await self._handle_before_agent_callback(ctx)
349+
before_callback_completed = True
350+
if event:
346351
yield event
347352
if ctx.end_invocation:
348353
return
@@ -354,8 +359,25 @@ async def _run() -> AsyncGenerator[Event, None]:
354359
if ctx.end_invocation:
355360
return
356361

362+
after_callback_called = True
357363
if event := await self._handle_after_agent_callback(ctx):
358364
yield event
365+
except asyncio.CancelledError:
366+
if (
367+
before_callback_completed
368+
and not after_callback_called
369+
and not ctx.end_invocation
370+
):
371+
try:
372+
await self._handle_after_agent_callback(ctx)
373+
except asyncio.CancelledError:
374+
raise
375+
except Exception: # pylint: disable=broad-except
376+
logger.exception(
377+
'after_agent_callback raised on cancellation;'
378+
' suppressing so original cancellation propagates.'
379+
)
380+
raise
359381
except Exception as e:
360382
await self._handle_agent_error_callback(ctx, e)
361383
raise
@@ -403,8 +425,12 @@ async def run_live(
403425
async def _run() -> AsyncGenerator[Event, None]:
404426
ctx = self._create_invocation_context(parent_context)
405427
async with _instrumentation.record_agent_invocation(ctx, self):
428+
before_callback_completed = False
429+
after_callback_called = False
406430
try:
407-
if event := await self._handle_before_agent_callback(ctx):
431+
event = await self._handle_before_agent_callback(ctx)
432+
before_callback_completed = True
433+
if event:
408434
yield event
409435
if ctx.end_invocation:
410436
return
@@ -413,8 +439,28 @@ async def _run() -> AsyncGenerator[Event, None]:
413439
async for event in agen:
414440
yield event
415441

442+
if ctx.end_invocation:
443+
return
444+
445+
after_callback_called = True
416446
if event := await self._handle_after_agent_callback(ctx):
417447
yield event
448+
except asyncio.CancelledError:
449+
if (
450+
before_callback_completed
451+
and not after_callback_called
452+
and not ctx.end_invocation
453+
):
454+
try:
455+
await self._handle_after_agent_callback(ctx)
456+
except asyncio.CancelledError:
457+
raise
458+
except Exception: # pylint: disable=broad-except
459+
logger.exception(
460+
'after_agent_callback raised on cancellation;'
461+
' suppressing so original cancellation propagates.'
462+
)
463+
raise
418464
except Exception as e:
419465
await self._handle_agent_error_callback(ctx, e)
420466
raise

0 commit comments

Comments
 (0)