From a3410fab7b25cc0e9c5908e23a087b501466df76 Mon Sep 17 00:00:00 2001 From: pguerra-ce Date: Wed, 3 Sep 2025 14:12:47 -0700 Subject: [PATCH] fix: Pass state_delta to runner in /run endpoint MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Merge https://github.com/google/adk-python/pull/2791 Fixes #2789 ## Summary Forward `state_delta` from the FastAPI `/run` request to `Runner.run_async(...)`, aligning behavior with the documented API and the `/run_sse` endpoint. ## Why The documentation for `/run` explicitly includes: > `state_delta` (object, optional): A delta of the state to apply before the run. However, the non‑SSE `/run` handler did not pass this value through, so `Runner.run_async` always received `None`. The `/run_sse` path already forwarded it correctly. ## Changes - `src/google/adk/cli/adk_web_server.py` - Add `state_delta=req.state_delta` to the "/run" handler’s `runner.run_async(...)` call. - `tests/unittests/cli/test_fast_api.py` - Add `test_agent_run_passes_state_delta` to test the fix. COPYBARA_INTEGRATE_REVIEW=https://github.com/google/adk-python/pull/2791 from pguerra-ce:fix-state-delta-missing-in-run 83eec8d28b80757e24ae900285eb59530863adbd PiperOrigin-RevId: 802703072 --- src/google/adk/cli/adk_web_server.py | 1 + tests/unittests/cli/test_fast_api.py | 37 ++++++++++++++++++++++++++++ 2 files changed, 38 insertions(+) diff --git a/src/google/adk/cli/adk_web_server.py b/src/google/adk/cli/adk_web_server.py index eaca7da2..1e0c3213 100644 --- a/src/google/adk/cli/adk_web_server.py +++ b/src/google/adk/cli/adk_web_server.py @@ -1020,6 +1020,7 @@ class AdkWebServer: user_id=req.user_id, session_id=req.session_id, new_message=req.new_message, + state_delta=req.state_delta, ) ) as agen: events = [event async for event in agen] diff --git a/tests/unittests/cli/test_fast_api.py b/tests/unittests/cli/test_fast_api.py index 423581df..b05bf3c4 100755 --- a/tests/unittests/cli/test_fast_api.py +++ b/tests/unittests/cli/test_fast_api.py @@ -34,6 +34,7 @@ from google.adk.evaluation.eval_result import EvalSetResult from google.adk.evaluation.eval_set import EvalSet from google.adk.evaluation.in_memory_eval_sets_manager import InMemoryEvalSetsManager from google.adk.events.event import Event +from google.adk.events.event_actions import EventActions from google.adk.runners import Runner from google.adk.sessions.base_session_service import ListSessionsResponse from google.genai import types @@ -94,6 +95,14 @@ def _event_3(): ) +def _event_state_delta(state_delta: dict[str, Any]): + return Event( + author="dummy agent", + invocation_id="invocation_id", + actions=EventActions(state_delta=state_delta), + ) + + # Define mocked async generator functions for the Runner async def dummy_run_live(self, session, live_request_queue): yield _event_1() @@ -110,6 +119,7 @@ async def dummy_run_async( user_id, session_id, new_message, + state_delta=None, run_config: RunConfig = RunConfig(), ): yield _event_1() @@ -119,6 +129,10 @@ async def dummy_run_async( await asyncio.sleep(0) yield _event_3() + await asyncio.sleep(0) + + if state_delta is not None: + yield _event_state_delta(state_delta) # Define a local mock for EvalCaseResult specific to fast_api tests @@ -744,6 +758,29 @@ def test_agent_run(test_app, create_test_session): logger.info("Agent run test completed successfully") +def test_agent_run_passes_state_delta(test_app, create_test_session): + """Test /run forwards state_delta and surfaces it in events.""" + info = create_test_session + payload = { + "app_name": info["app_name"], + "user_id": info["user_id"], + "session_id": info["session_id"], + "new_message": {"role": "user", "parts": [{"text": "Hello"}]}, + "streaming": False, + "state_delta": {"k": "v", "count": 1}, + } + + # Verify the response + response = test_app.post("/run", json=payload) + assert response.status_code == 200 + data = response.json() + assert isinstance(data, list) + assert len(data) == 4 + + # Verify we got the expected event + assert data[3]["actions"]["stateDelta"] == payload["state_delta"] + + def test_list_artifact_names(test_app, create_test_session): """Test listing artifact names for a session.""" info = create_test_session