Skip to content

Commit 6de43b0

Browse files
lwangverizoncopybara-github
authored andcommitted
fix: apply state_delta when resuming without a new_message
Merge #6645 Fixes #6644 PiperOrigin-RevId: 978324019
1 parent bfc46fe commit 6de43b0

3 files changed

Lines changed: 277 additions & 0 deletions

File tree

‎src/google/adk/runners.py‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -702,6 +702,40 @@ async def _append_user_event(
702702
session=ic.session, event=event
703703
)
704704

705+
async def _append_state_delta_event(
706+
self,
707+
ic: InvocationContext,
708+
state_delta: dict[str, Any],
709+
) -> Event:
710+
"""Appends an event to the session carrying only a state delta.
711+
712+
Used when resuming an invocation without a new message, so that any
713+
caller-supplied state delta is still persisted to the session rather than
714+
being dropped because there is no new message event to attach it to.
715+
716+
Args:
717+
ic: The invocation context for the run.
718+
state_delta: The state delta dictionary to append.
719+
720+
Returns:
721+
The appended event, matching the return convention of
722+
the user message event.
723+
"""
724+
event = Event(
725+
invocation_id=ic.invocation_id,
726+
author='user',
727+
actions=EventActions(state_delta=state_delta),
728+
)
729+
if event.isolation_scope is None:
730+
active_scope = _find_active_task_scope(ic.session)
731+
if active_scope is not None:
732+
event.isolation_scope, _ = active_scope
733+
_apply_run_config_custom_metadata(event, ic.run_config)
734+
ic.stamp_event_branch_context(event)
735+
return await self.session_service.append_event(
736+
session=ic.session, event=event
737+
)
738+
705739
def _find_user_message_for_invocation(
706740
self, events: list[Event], invocation_id: str
707741
) -> types.Content | None:
@@ -1913,6 +1947,11 @@ async def _setup_context_for_resumed_invocation(
19131947
run_config=run_config,
19141948
state_delta=state_delta,
19151949
)
1950+
elif state_delta:
1951+
# Resuming without a new message: there is no user message event to
1952+
# carry the delta, so append it as a content-less event instead of
1953+
# dropping it.
1954+
await self._append_state_delta_event(invocation_context, state_delta)
19161955
# Step 4: Populate agent states for the current invocation.
19171956
invocation_context.populate_invocation_agent_states()
19181957
# Step 5: Set agent to run for the invocation.

‎src/google/adk/workflow/_node_runner_utils.py‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -175,6 +175,15 @@ async def _run() -> AsyncGenerator[Event, None]:
175175
)
176176
if yield_user_message and user_event:
177177
yield user_event
178+
elif state_delta:
179+
# Resuming without a new message: there is no user message event to
180+
# carry the delta, so append it as a content-less event instead of
181+
# dropping it.
182+
delta_event = await runner._append_state_delta_event( # pylint: disable=protected-access
183+
ic, state_delta
184+
)
185+
if yield_user_message and delta_event:
186+
yield delta_event
178187

179188
# Run before_run callbacks. A returned Content halts execution and ends
180189
# the run with that content (same contract as the non-workflow path).

‎tests/unittests/test_runners.py‎

Lines changed: 229 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -628,6 +628,235 @@ async def _run_async_impl(
628628
assert isinstance(exc_info.value.__cause__, asyncio.CancelledError)
629629

630630

631+
@pytest.mark.asyncio
632+
async def test_run_async_applies_state_delta_when_resuming_without_new_message():
633+
"""Resuming by invocation_id should still apply a caller-supplied delta."""
634+
635+
session_service = InMemorySessionService()
636+
runner = Runner(
637+
app_name=TEST_APP_ID,
638+
agent=MockAgent("test_agent"),
639+
session_service=session_service,
640+
artifact_service=InMemoryArtifactService(),
641+
auto_create_session=True,
642+
)
643+
runner.resumability_config = ResumabilityConfig(is_resumable=True)
644+
645+
# Seed the session with an invocation to resume.
646+
async with aclosing(
647+
runner.run_async(
648+
user_id=TEST_USER_ID,
649+
session_id=TEST_SESSION_ID,
650+
new_message=types.Content(
651+
role="user", parts=[types.Part(text="hello")]
652+
),
653+
)
654+
) as agen:
655+
async for _ in agen:
656+
pass
657+
658+
session = await session_service.get_session(
659+
app_name=TEST_APP_ID, user_id=TEST_USER_ID, session_id=TEST_SESSION_ID
660+
)
661+
invocation_id = session.events[0].invocation_id
662+
663+
state_delta = {"resumed_key": "resumed_value"}
664+
665+
async with aclosing(
666+
runner.run_async(
667+
user_id=TEST_USER_ID,
668+
session_id=TEST_SESSION_ID,
669+
invocation_id=invocation_id,
670+
state_delta=state_delta,
671+
)
672+
) as agen:
673+
async for _ in agen:
674+
pass
675+
676+
session = await session_service.get_session(
677+
app_name=TEST_APP_ID, user_id=TEST_USER_ID, session_id=TEST_SESSION_ID
678+
)
679+
680+
assert session.state["resumed_key"] == "resumed_value"
681+
682+
683+
@pytest.mark.asyncio
684+
async def test_run_async_applies_state_delta_when_resuming_without_new_message_llm_agent():
685+
"""Resuming by invocation_id should apply caller-supplied delta for LLM agent."""
686+
687+
session_service = InMemorySessionService()
688+
runner = Runner(
689+
app_name=TEST_APP_ID,
690+
agent=MockLlmAgent("test_llm_agent"),
691+
session_service=session_service,
692+
artifact_service=InMemoryArtifactService(),
693+
auto_create_session=True,
694+
)
695+
runner.resumability_config = ResumabilityConfig(is_resumable=True)
696+
697+
# Seed the session with an invocation to resume.
698+
async with aclosing(
699+
runner.run_async(
700+
user_id=TEST_USER_ID,
701+
session_id=TEST_SESSION_ID,
702+
new_message=types.Content(
703+
role="user", parts=[types.Part(text="hello")]
704+
),
705+
)
706+
) as agen:
707+
async for _ in agen:
708+
pass
709+
710+
session = await session_service.get_session(
711+
app_name=TEST_APP_ID, user_id=TEST_USER_ID, session_id=TEST_SESSION_ID
712+
)
713+
invocation_id = session.events[0].invocation_id
714+
715+
state_delta = {"resumed_key": "resumed_value"}
716+
717+
async with aclosing(
718+
runner.run_async(
719+
user_id=TEST_USER_ID,
720+
session_id=TEST_SESSION_ID,
721+
invocation_id=invocation_id,
722+
state_delta=state_delta,
723+
)
724+
) as agen:
725+
async for _ in agen:
726+
pass
727+
728+
session = await session_service.get_session(
729+
app_name=TEST_APP_ID, user_id=TEST_USER_ID, session_id=TEST_SESSION_ID
730+
)
731+
732+
assert session.state["resumed_key"] == "resumed_value"
733+
delta_event = [
734+
e
735+
for e in session.events
736+
if e.actions and e.actions.state_delta and not e.content
737+
][0]
738+
assert delta_event.branch is None
739+
740+
741+
@pytest.mark.asyncio
742+
async def test_run_node_async_yields_state_delta_when_resuming_with_yield_user_message():
743+
"""run_node_async yields delta event when yield_user_message=True on resume."""
744+
from typing import Any
745+
746+
from google.adk.agents.context import Context
747+
from google.adk.workflow import _node_runner_utils
748+
from google.adk.workflow._base_node import BaseNode
749+
750+
class _TestNode(BaseNode):
751+
752+
async def _run_impl(
753+
self, *, ctx: Context, node_input: Any
754+
) -> AsyncGenerator[Event, None]:
755+
yield Event(
756+
author=self.name,
757+
content=types.Content(
758+
role="model", parts=[types.Part(text="node response")]
759+
),
760+
)
761+
762+
session_service = InMemorySessionService()
763+
node = _TestNode(name="test_node")
764+
runner = Runner(
765+
app_name=TEST_APP_ID,
766+
agent=MockAgent("test_agent"),
767+
session_service=session_service,
768+
artifact_service=InMemoryArtifactService(),
769+
auto_create_session=True,
770+
)
771+
session = await session_service.create_session(
772+
app_name=TEST_APP_ID, user_id=TEST_USER_ID, session_id=TEST_SESSION_ID
773+
)
774+
seed_event = Event(
775+
invocation_id="inv_1",
776+
author="user",
777+
content=types.Content(role="user", parts=[types.Part(text="initial")]),
778+
)
779+
await session_service.append_event(session, seed_event)
780+
781+
state_delta = {"resumed_key": "resumed_value"}
782+
events = []
783+
async for event in _node_runner_utils.run_node_async(
784+
runner,
785+
user_id=TEST_USER_ID,
786+
session_id=TEST_SESSION_ID,
787+
invocation_id="inv_1",
788+
state_delta=state_delta,
789+
yield_user_message=True,
790+
node=node,
791+
session=session,
792+
):
793+
events.append(event)
794+
795+
user_events = [e for e in events if e.author == "user"]
796+
assert len(user_events) == 1
797+
assert user_events[0].actions.state_delta == state_delta
798+
assert session.state["resumed_key"] == "resumed_value"
799+
800+
801+
@pytest.mark.asyncio
802+
async def test_run_node_async_does_not_yield_state_delta_when_resuming_without_yield_user_message():
803+
"""run_node_async does not yield delta event by default on resume."""
804+
from typing import Any
805+
806+
from google.adk.agents.context import Context
807+
from google.adk.workflow import _node_runner_utils
808+
from google.adk.workflow._base_node import BaseNode
809+
810+
class _TestNode(BaseNode):
811+
812+
async def _run_impl(
813+
self, *, ctx: Context, node_input: Any
814+
) -> AsyncGenerator[Event, None]:
815+
yield Event(
816+
author=self.name,
817+
content=types.Content(
818+
role="model", parts=[types.Part(text="node response")]
819+
),
820+
)
821+
822+
session_service = InMemorySessionService()
823+
node = _TestNode(name="test_node")
824+
runner = Runner(
825+
app_name=TEST_APP_ID,
826+
agent=MockAgent("test_agent"),
827+
session_service=session_service,
828+
artifact_service=InMemoryArtifactService(),
829+
auto_create_session=True,
830+
)
831+
session = await session_service.create_session(
832+
app_name=TEST_APP_ID, user_id=TEST_USER_ID, session_id=TEST_SESSION_ID
833+
)
834+
seed_event = Event(
835+
invocation_id="inv_1",
836+
author="user",
837+
content=types.Content(role="user", parts=[types.Part(text="initial")]),
838+
)
839+
await session_service.append_event(session, seed_event)
840+
841+
state_delta = {"resumed_key": "resumed_value"}
842+
events = []
843+
async for event in _node_runner_utils.run_node_async(
844+
runner,
845+
user_id=TEST_USER_ID,
846+
session_id=TEST_SESSION_ID,
847+
invocation_id="inv_1",
848+
state_delta=state_delta,
849+
yield_user_message=False,
850+
node=node,
851+
session=session,
852+
):
853+
events.append(event)
854+
855+
user_events = [e for e in events if e.author == "user"]
856+
assert not user_events
857+
assert session.state["resumed_key"] == "resumed_value"
858+
859+
631860
@pytest.mark.asyncio
632861
async def test_run_async_propagates_invocation_id():
633862
"""run_async should propagate invocation_id to the invocation context and events."""

0 commit comments

Comments
 (0)