From 32922e10eff468faa4ccd5c3222c48f1eb50d66b Mon Sep 17 00:00:00 2001 From: gioboa Date: Thu, 27 Aug 2026 20:52:48 +0200 Subject: [PATCH] fix(cli): prevent SSE disconnect corruption --- src/google/adk/cli/api_server.py | 176 +++++++++++++++------------ tests/unittests/cli/test_fast_api.py | 57 ++++++++- 2 files changed, 155 insertions(+), 78 deletions(-) diff --git a/src/google/adk/cli/api_server.py b/src/google/adk/cli/api_server.py index 80143817688..bd983a55dd4 100644 --- a/src/google/adk/cli/api_server.py +++ b/src/google/adk/cli/api_server.py @@ -855,6 +855,7 @@ def __init__( self.runners_to_clean: set[str] = set() self.current_app_name_ref: SharedValue[str] = SharedValue(value="") self.runner_dict: dict[str, Runner] = {} + self._run_tasks: set[asyncio.Task[None]] = set() self.url_prefix = url_prefix self.auto_create_session = auto_create_session self.trigger_sources = trigger_sources @@ -1164,6 +1165,10 @@ async def internal_lifespan(app: FastAPI): yield finally: tear_down_observer(observer, self) + run_tasks = list(self._run_tasks) + for task in run_tasks: + task.cancel() + await asyncio.gather(*run_tasks, return_exceptions=True) # Create tasks for all runner closures to run concurrently await cleanup.close_runners(list(self.runner_dict.values())) @@ -1895,86 +1900,103 @@ async def run_agent_sse(req: RunAgentRequest) -> StreamingResponse: # Convert the events to properly formatted SSE async def event_generator(): - is_closing = False - original_exc = None + client_connected = asyncio.Event() + client_connected.set() + event_queue: asyncio.Queue[ + tuple[Optional[Event], Optional[Exception]] + ] = asyncio.Queue(maxsize=1) + + async def run_agent() -> None: + 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, + run_config=RunConfig( + streaming_mode=stream_mode, + custom_metadata=req.custom_metadata, + ), + invocation_id=req.invocation_id, + ) + ) as agen: + async for event in agen: + if client_connected.is_set(): + await event_queue.put((event, None)) + except asyncio.CancelledError: + client_connected.clear() + raise + except Exception as e: + if client_connected.is_set(): + await event_queue.put((None, e)) + else: + logger.exception("Detached agent run failed: %s", e) + finally: + if client_connected.is_set(): + await event_queue.put((None, None)) + + run_task = asyncio.create_task(run_agent()) + self._run_tasks.add(run_task) + run_task.add_done_callback(self._run_tasks.discard) 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, - run_config=RunConfig( - streaming_mode=stream_mode, - custom_metadata=req.custom_metadata, - ), - invocation_id=req.invocation_id, + while True: + event, error = await event_queue.get() + if error is not None: + logger.exception( + "Error in event_generator: %s", + error, + exc_info=(type(error), error, error.__traceback__), ) - ) as agen: - try: - async for event in agen: - # ADK Web renders artifacts from `actions.artifactDelta` - # during part processing *and* during action processing - # 1) the original event with `artifactDelta` cleared (content) - # 2) a content-less "action-only" event carrying `artifactDelta` - events_to_stream = [event] - if ( - not req.function_call_event_id - and event.actions.artifact_delta - and event.content - and event.content.parts - ): - content_event = event.model_copy(deep=True) - content_event.actions.artifact_delta = {} - artifact_event = event.model_copy(deep=True) - artifact_event.content = None - events_to_stream = [content_event, artifact_event] - - for event_to_stream in events_to_stream: - sse_event = event_to_stream.model_dump_json( - exclude_none=True, - by_alias=True, - ) - logger.debug( - "Generated event in agent run streaming: %s", sse_event - ) - yield f"data: {sse_event}\n\n" - except (GeneratorExit, asyncio.CancelledError) as e: - is_closing = True - original_exc = e - raise - except Exception as e: - original_exc = e - raise - except Exception as e: - if original_exc: - if e is not original_exc: - logger.exception("Error during generator cleanup: %s", e) - if is_closing: - raise original_exc from e - logger.exception("Error in event_generator: %s", original_exc) - error_details = { - "error_type": type(original_exc).__name__, - "error_message": str(original_exc), - "timestamp": time.time(), - } - if logger.isEnabledFor(logging.DEBUG): - error_details["stacktrace"] = "".join( - traceback.format_exception( - type(original_exc), - original_exc, - original_exc.__traceback__, - ) + error_details = { + "error_type": type(error).__name__, + "error_message": str(error), + "timestamp": time.time(), + } + if logger.isEnabledFor(logging.DEBUG): + error_details["stacktrace"] = "".join( + traceback.format_exception( + type(error), error, error.__traceback__ + ) + ) + yield ( + "data:" + f" {json.dumps({'error': f'{type(error).__name__}: {error}', 'error_details': error_details})}\n\n" ) - yield ( - "data:" - f" {json.dumps({'error': f'{type(original_exc).__name__}: {original_exc}', 'error_details': error_details})}\n\n" - ) - return - logger.exception( - "Error during generator cleanup after completion: %s", e - ) - raise e + return + if event is None: + return + + # ADK Web renders artifacts from `actions.artifactDelta` + # during part processing *and* during action processing: + # 1) the original event with `artifactDelta` cleared (content) + # 2) a content-less "action-only" event carrying `artifactDelta` + events_to_stream = [event] + if ( + not req.function_call_event_id + and event.actions.artifact_delta + and event.content + and event.content.parts + ): + content_event = event.model_copy(deep=True) + content_event.actions.artifact_delta = {} + artifact_event = event.model_copy(deep=True) + artifact_event.content = None + events_to_stream = [content_event, artifact_event] + + for event_to_stream in events_to_stream: + sse_event = event_to_stream.model_dump_json( + exclude_none=True, + by_alias=True, + ) + logger.debug( + "Generated event in agent run streaming: %s", sse_event + ) + yield f"data: {sse_event}\n\n" + finally: + client_connected.clear() + while not event_queue.empty(): + event_queue.get_nowait() # Returns a streaming response with the proper media type for SSE return StreamingResponse( diff --git a/tests/unittests/cli/test_fast_api.py b/tests/unittests/cli/test_fast_api.py index a3fe28d35a3..745db4d7a9d 100644 --- a/tests/unittests/cli/test_fast_api.py +++ b/tests/unittests/cli/test_fast_api.py @@ -2116,6 +2116,8 @@ async def test_agent_run_sse_disconnect_with_cleanup_exception_and_cancellation( from google.adk.cli.api_server import RunAgentRequest info = create_test_session + release_run = asyncio.Event() + cleanup_finished = asyncio.Event() class MockAsyncGenerator: @@ -2136,10 +2138,11 @@ async def __anext__(self): ), ) # Block indefinitely to allow cancellation simulation - await asyncio.sleep(10) + await release_run.wait() raise StopAsyncIteration async def aclose(self): + cleanup_finished.set() raise ValueError("cleanup failed") def run_async_mock(self, **kwargs): @@ -2190,6 +2193,58 @@ def run_async_mock(self, **kwargs): with pytest.raises(asyncio.CancelledError): await task + release_run.set() + await asyncio.wait_for(cleanup_finished.wait(), timeout=1) + + +async def test_agent_run_sse_disconnect_keeps_active_run_alive( + test_app, create_test_session, monkeypatch +): + """Disconnecting an SSE client leaves its active agent run running.""" + from google.adk.cli.api_server import RunAgentRequest + + info = create_test_session + continue_run = asyncio.Event() + run_finished = asyncio.Event() + was_cancelled = asyncio.Event() + + async def run_async_mock(self, **kwargs): + del self, kwargs + try: + yield _event_1() + await continue_run.wait() + yield _event_2() + except asyncio.CancelledError: + was_cancelled.set() + raise + finally: + run_finished.set() + + monkeypatch.setattr(Runner, "run_async", run_async_mock) + + handler = next( + route.endpoint + for route in test_app.app.routes + if route.path == "/run_sse" + ) + req = RunAgentRequest( + app_name=info["app_name"], + user_id=info["user_id"], + session_id=info["session_id"], + new_message={"role": "user", "parts": [{"text": "Hello agent"}]}, + streaming=True, + ) + + response = await handler(req) + generator = response.body_iterator + assert "LLM reply" in await generator.__anext__() + + await generator.aclose() + continue_run.set() + await asyncio.wait_for(run_finished.wait(), timeout=1) + + assert not was_cancelled.is_set() + def test_list_artifact_names(test_app, create_test_session): """Test listing artifact names for a session."""