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
7 changes: 6 additions & 1 deletion src/google/adk/cli/api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,7 @@
from ..models._service_tier import ServiceTier
from ..plugins.base_plugin import BasePlugin
from ..runners import Runner
from ..runners import _CALLER_CLOSED_EARLY_MSG
from ..sessions.base_session_service import BaseSessionService
from ..sessions.session import Session
from ..utils._telemetry_config import read_telemetry_consent
Expand Down Expand Up @@ -2018,7 +2019,11 @@ async def monitor():
"Client disconnected. Aborting agent run for session %s.",
req.session_id,
)
worker_task.cancel()
# Cancel with the caller-closed-early marker instead of a bare
# cancel: a disconnect is a caller that stopped consuming events,
# so plugins still get after_run_callback. A bare cancellation is
# treated as an external cancel and deliberately skips it.
worker_task.cancel(_CALLER_CLOSED_EARLY_MSG)
break
except asyncio.CancelledError:
pass
Expand Down
85 changes: 85 additions & 0 deletions tests/unittests/cli/test_fast_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -5932,6 +5932,91 @@ async def mock_receive():
assert was_cancelled["value"] is True


def test_agent_run_disconnect_marks_cancellation_as_caller_closed_early(
test_app, create_test_session, monkeypatch
):
"""Test that /run cancels the run task with the caller-closed-early marker.

A disconnect means the caller stopped consuming events, which is exactly what
the marker means, so plugins still receive after_run_callback. A bare
cancellation counts as an external cancel and deliberately skips it.
"""
import starlette.requests

from google.adk.runners import _CALLER_CLOSED_EARLY_MSG

info = create_test_session
trigger_disconnect: dict[str, bool] = {"value": False}
cancel_args: list[tuple] = []

async def run_async_mock(
self,
*,
user_id: str,
session_id: str,
invocation_id: Optional[str] = None,
new_message: Optional[types.Content] = None,
state_delta: Optional[dict[str, Any]] = None,
run_config: Optional[RunConfig] = None,
):
del (
self,
user_id,
session_id,
invocation_id,
new_message,
state_delta,
run_config,
)
try:
yield _event_1()
trigger_disconnect["value"] = True
await asyncio.sleep(1.0)
yield _event_2()
except asyncio.CancelledError as e:
cancel_args.append(e.args)
raise

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

# Monkeypatch starlette.requests.Request.__init__ to inject simulated disconnect
original_init = starlette.requests.Request.__init__

def custom_init(self, *args, **kwargs):
original_init(self, *args, **kwargs)
original_receive = self._receive
call_count = 0

async def mock_receive():
nonlocal call_count
call_count += 1
if call_count == 1:
return await original_receive()

while not trigger_disconnect["value"]:
await asyncio.sleep(0.01)
return {"type": "http.disconnect"}

self._receive = mock_receive
self.__dict__["receive"] = mock_receive

monkeypatch.setattr(starlette.requests.Request, "__init__", custom_init)

payload = {
"app_name": info["app_name"],
"user_id": info["user_id"],
"session_id": info["session_id"],
"new_message": {"role": "user", "parts": [{"text": "Hello agent"}]},
"streaming": False,
}

response = test_app.post("/run", json=payload)

assert response.status_code == 499
assert cancel_args, "the run task was never cancelled"
assert cancel_args[0] == (_CALLER_CLOSED_EARLY_MSG,)


#################################################
# Gemini Enterprise Tests
#################################################
Expand Down
Loading