Skip to content
Closed
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
31 changes: 21 additions & 10 deletions src/google/adk/cli/api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -1991,18 +1991,28 @@ async def run_agent(req: RunAgentRequest, request: Request) -> Response:
),
)

abort_signal = asyncio.Event()

async def worker():
run_async_kwargs: dict[str, Any] = {
"user_id": req.user_id,
"session_id": req.session_id,
"new_message": req.new_message,
"state_delta": req.state_delta,
"invocation_id": req.invocation_id,
"run_config": run_config,
}
try:
async with Aclosing(
runner.run_async(
user_id=req.user_id,
session_id=req.session_id,
new_message=req.new_message,
state_delta=req.state_delta,
invocation_id=req.invocation_id,
run_config=run_config,
)
) as agen:
params = inspect.signature(runner.run_async).parameters
if "abort_signal" in params or any(
p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values()
):
run_async_kwargs["abort_signal"] = abort_signal
except (ValueError, TypeError):
run_async_kwargs["abort_signal"] = abort_signal

try:
async with Aclosing(runner.run_async(**run_async_kwargs)) as agen:
return [public_event(event) async for event in agen]
except SessionNotFoundError as e:
raise HTTPException(status_code=404, detail=str(e)) from e
Expand All @@ -2018,6 +2028,7 @@ async def monitor():
"Client disconnected. Aborting agent run for session %s.",
req.session_id,
)
abort_signal.set()
worker_task.cancel()
break
except asyncio.CancelledError:
Expand Down
4 changes: 3 additions & 1 deletion src/google/adk/runners.py
Original file line number Diff line number Diff line change
Expand Up @@ -1645,7 +1645,9 @@ async def _exec_with_plugin(
await _notify_run_error(plugin_manager, invocation_context, e)
raise
except asyncio.CancelledError as e:
if e.args and e.args[0] == _CALLER_CLOSED_EARLY_MSG:
if (
e.args and e.args[0] == _CALLER_CLOSED_EARLY_MSG
) or invocation_context.is_aborted:
closing_early = True
else:
run_error = e
Expand Down
9 changes: 6 additions & 3 deletions src/google/adk/workflow/_node_runner_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -302,7 +302,9 @@ async def _drive_root_node() -> None:
await _notify_run_error(ic.plugin_manager, ic, e)
raise
except asyncio.CancelledError as e:
if e.args and e.args[0] == _CALLER_CLOSED_EARLY_MSG:
if (
e.args and e.args[0] == _CALLER_CLOSED_EARLY_MSG
) or ic.is_aborted:
closing_early = True
else:
run_error = e
Expand All @@ -313,8 +315,9 @@ async def _drive_root_node() -> None:
run_error = e
raise
finally:
# Success path (also caller early-stop via GeneratorExit or
# _CALLER_CLOSED_EARLY_MSG): run after_run and compaction.
# Success path (also caller early-stop via GeneratorExit,
# _CALLER_CLOSED_EARLY_MSG, or abort_signal): run after_run and
# compaction.
# _cleanup_root_task has 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
Expand Down
132 changes: 131 additions & 1 deletion tests/unittests/cli/test_fast_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -3862,9 +3862,25 @@ async def _run_async_impl(self, invocation_context):
tool_in_flight.set()
await asyncio.sleep(5.0)

from google.adk.apps.app import App
from google.adk.plugins.base_plugin import BasePlugin

after_run_called = False

class _AfterRunPlugin(BasePlugin):

async def after_run_callback(self, *, invocation_context):
nonlocal after_run_called
after_run_called = True

slow_agent = SlowToolAgent("slow_tool_agent")
loaded_app = App(
name=info["app_name"],
root_agent=slow_agent,
plugins=[_AfterRunPlugin(name="after_run")],
)
monkeypatch.setattr(
mock_agent_loader, "load_agent", lambda app_name: slow_agent
mock_agent_loader, "load_agent", lambda app_name: loaded_app
)

client = _create_test_client(
Expand Down Expand Up @@ -3913,6 +3929,7 @@ async def send(message):
assert any("slow_tool" in chunk for chunk in sent_chunks)
assert len(captured_contexts) == 1
assert captured_contexts[0].is_aborted is True
assert after_run_called is True

# Verify the dangling FunctionCall was sealed with a synthetic FunctionResponse in session
session = await mock_session_service.get_session(
Expand Down Expand Up @@ -5932,6 +5949,119 @@ async def mock_receive():
assert was_cancelled["value"] is True


async def test_agent_run_disconnect_seals_dangling_call_and_runs_after_run(
create_test_session,
mock_session_service,
mock_agent_loader,
mock_eval_sets_manager,
mock_eval_set_results_manager,
monkeypatch,
):
"""Test that /run disconnect seals dangling FunctionCall and runs after_run."""
from google.adk.apps.app import App
from google.adk.plugins.base_plugin import BasePlugin
import starlette.requests

info = create_test_session
captured_contexts = []
tool_in_flight = asyncio.Event()
after_run_called = False

monkeypatch.setattr(Runner, "run_async", _ORIGINAL_RUNNER_RUN_ASYNC)

class _AfterRunPlugin(BasePlugin):

async def after_run_callback(self, *, invocation_context):
nonlocal after_run_called
after_run_called = True

class SlowToolAgent(BaseAgent):

def __init__(self, name: str):
super().__init__(name=name, sub_agents=[])

async def _run_async_impl(self, invocation_context):
captured_contexts.append(invocation_context)
fc = types.Part.from_function_call(name="slow_tool", args={"q": "test"})
fc.function_call.id = "call_run_1"
yield Event(
invocation_id=invocation_context.invocation_id,
author=self.name,
content=types.Content(role="model", parts=[fc]),
)
tool_in_flight.set()
await asyncio.sleep(5.0)

slow_agent = SlowToolAgent("slow_tool_agent")
loaded_app = App(
name=info["app_name"],
root_agent=slow_agent,
plugins=[_AfterRunPlugin(name="after_run")],
)
monkeypatch.setattr(
mock_agent_loader, "load_agent", lambda app_name: loaded_app
)

client = _create_test_client(
mock_session_service,
InMemoryArtifactService(),
InMemoryMemoryService(),
mock_agent_loader,
mock_eval_sets_manager,
mock_eval_set_results_manager,
)
app = client.app
handler = None
for route in app.routes:
if route.path == "/run":
handler = route.endpoint
break
assert handler is not None

req = RunAgentRequest(
app_name=info["app_name"],
user_id=info["user_id"],
session_id=info["session_id"],
new_message={"role": "user", "parts": [{"text": "Run slow tool"}]},
streaming=False,
)

async def receive():
await tool_in_flight.wait()
return {"type": "http.disconnect"}

request = starlette.requests.Request(
{
"type": "http",
"method": "POST",
"path": "/run",
"headers": [],
"asgi": {"spec_version": "2.1"},
},
receive=receive,
)

response = await handler(req, request)
assert response.status_code == 499
assert len(captured_contexts) == 1
assert captured_contexts[0].is_aborted is True
assert after_run_called is True

session = await mock_session_service.get_session(
app_name=info["app_name"],
user_id=info["user_id"],
session_id=info["session_id"],
)
abort_events = [
e for e in session.events if e.error_code == "INVOCATION_ABORTED"
]
assert len(abort_events) == 1
frs = abort_events[0].get_function_responses()
assert len(frs) == 1
assert frs[0].id == "call_run_1"
assert frs[0].name == "slow_tool"


#################################################
# Gemini Enterprise Tests
#################################################
Expand Down
54 changes: 54 additions & 0 deletions tests/unittests/test_runners.py
Original file line number Diff line number Diff line change
Expand Up @@ -5516,6 +5516,60 @@ async def failing_append(session, event):
assert len(after_run_calls) == 1


@pytest.mark.asyncio
@pytest.mark.parametrize(
"make_agent", [_legacy_tool_agent, _llm_tool_agent], ids=["legacy", "llm"]
)
async def test_run_async_aborted_and_cancelled_executes_after_run_plugin(
make_agent,
):
"""Cancelling an aborted run seals dangling calls and runs after_run."""
after_run_calls = []

class _AfterRunPlugin(BasePlugin):

async def after_run_callback(self, *, invocation_context):
after_run_calls.append(invocation_context.invocation_id)

runner = _abort_runner(
make_agent(),
plugins=[_AfterRunPlugin(name="after_run")],
)
abort_signal = asyncio.Event()
tool_started = asyncio.Event()

async def _consume():
async with aclosing(
runner.run_async(
user_id=TEST_USER_ID,
session_id="s",
new_message=types.Content(
role="user", parts=[types.Part(text="Run")]
),
abort_signal=abort_signal,
)
) as agen:
async for event in agen:
if _has_fc(event):
tool_started.set()

task = asyncio.create_task(_consume())
await tool_started.wait()
abort_signal.set()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task

session = await runner.session_service.get_session(
app_name=TEST_APP_ID, user_id=TEST_USER_ID, session_id="s"
)
sealed = _abort_events(session.events)
assert [fr.name for e in sealed for fr in e.get_function_responses()] == [
"_slow_tool"
]
assert len(after_run_calls) == 1


@pytest.mark.asyncio
async def test_run_async_aborted_task_sub_agent_does_not_capture_next_turn():
"""Aborting inside a task sub-agent seals both scopes; the next turn reaches the coordinator."""
Expand Down
Loading