From 3f0ac0ba31ed0bb88d5a708326c3e539e5357548 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Wed, 19 Aug 2026 23:47:07 +0530 Subject: [PATCH 1/3] fix(python): separate provider input from UI snapshot in AG-UI service-session mode --- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 21 ++++- .../ag-ui/tests/ag_ui/test_snapshots.py | 82 ++++++++++++++++++- 2 files changed, 100 insertions(+), 3 deletions(-) diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py index fb871cdb95b..1e84ef6d7bd 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py @@ -2268,9 +2268,12 @@ async def run_agent_stream( if stored_snapshot is not None: if resume_payload is not None and stored_pending_approval_interrupt_ids: - raw_messages = snapshot_session.resume_seeded_messages(raw_messages) seeded_resume_from_snapshot = True - else: + + if not config.use_service_session: + raw_messages = snapshot_session.resume_seeded_messages(raw_messages) + + elif not config.use_service_session: raw_messages = _reconstruct_messages_from_thread_snapshot( stored_messages=stored_snapshot.messages, incoming_messages=raw_messages, @@ -2378,6 +2381,15 @@ async def run_agent_stream( protected_tool_call_ids=protected_tool_call_ids, ) + if config.use_service_session and stored_snapshot is not None: + if seeded_resume_from_snapshot: + snapshot_messages = snapshot_session.resume_seeded_messages(snapshot_messages) + else: + snapshot_messages = _reconstruct_messages_from_thread_snapshot( + stored_messages=stored_snapshot.messages, + incoming_messages=snapshot_messages, + stored_interrupt=stored_snapshot.interrupt, + ) # Check for structured output mode (skip text content) skip_text = False response_format: type[Any] | None = None @@ -2840,6 +2852,11 @@ async def run_agent_stream( flow.pending_tool_calls or flow.tool_results or flow.accumulated_text or flow.reasoning_messages ) latest_messages_snapshot = snapshot_messages + + if config.use_service_session and stored_snapshot is not None and not should_emit_snapshot: + latest_messages_snapshot = _event_messages_to_snapshot_dicts( + list(_build_messages_snapshot(flow, snapshot_messages).messages) + ) if should_emit_snapshot: # Always fold this turn's output into the persisted snapshot, even when the # outbound MESSAGES_SNAPSHOT event is suppressed for predictive tools. diff --git a/python/packages/ag-ui/tests/ag_ui/test_snapshots.py b/python/packages/ag-ui/tests/ag_ui/test_snapshots.py index e2ea85e37df..8ed1b596d21 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_snapshots.py +++ b/python/packages/ag-ui/tests/ag_ui/test_snapshots.py @@ -4,7 +4,15 @@ from dataclasses import fields -from agent_framework_ag_ui import AGUIThreadSnapshot, AGUIThreadSnapshotStore, InMemoryAGUIThreadSnapshotStore +import pytest +from agent_framework import AgentResponseUpdate, Content + +from agent_framework_ag_ui import ( + AgentFrameworkAgent, + AGUIThreadSnapshot, + AGUIThreadSnapshotStore, + InMemoryAGUIThreadSnapshotStore, +) def test_thread_snapshot_model_contains_replayable_and_private_snapshot_fields() -> None: @@ -204,3 +212,75 @@ async def test_in_memory_snapshot_store_rejects_invalid_keys() -> None: await store.delete(scope=None, thread_id="thread-1") # type: ignore[arg-type] # ty: ignore[invalid-argument-type] with pytest.raises(ValueError): await store.clear(scope="") + + +class _InputSpyAgent: + name = "spy" + description = "" + default_options: dict = {} + context_providers: list = [] + + def __init__(self) -> None: + self.calls: list[dict] = [] + + def run(self, messages, *, session, stream=False, **kwargs): + async def updates(): + self.calls.append( + { + "roles": [m.role for m in messages], + "service_session_id": session.service_session_id, + } + ) + yield AgentResponseUpdate( + contents=[Content.from_text("ACK")], + role="assistant", + response_id=f"resp-{len(self.calls)}", + ) + + return updates() + + +async def _drain(runner, body): + return [event async for event in runner.run(body)] + + +@pytest.mark.asyncio +async def test_service_session_snapshot_split_authority() -> None: + """Verify use_service_session + snapshot_store separates provider input from UI hydration. + + - Provider must receive ONLY incremental input + - Snapshot store must retain FULL history for UI hydration + """ + agent = _InputSpyAgent() + store = InMemoryAGUIThreadSnapshotStore() + runner = AgentFrameworkAgent( + agent=agent, + use_service_session=True, + snapshot_store=store, + ) + + first_turn = { + "threadId": "conv_FHA_SESSION", + "__ag_ui_snapshot_scope": "split-auth-test", + "messages": [{"id": "u1", "role": "user", "content": "first"}], + } + second_turn = { + "threadId": "conv_FHA_SESSION", + "__ag_ui_snapshot_scope": "split-auth-test", + "messages": [{"id": "u2", "role": "user", "content": "second"}], + } + + await _drain(runner, first_turn) + await _drain(runner, second_turn) + + assert agent.calls[1]["roles"] == ["user"], ( + f"Expected incremental-only input for service-session mode, got: {agent.calls[1]['roles']}" + ) + assert agent.calls[1]["service_session_id"] == "conv_FHA_SESSION" + + snapshot = await store.get(scope="split-auth-test", thread_id="conv_FHA_SESSION") + assert snapshot is not None, "Snapshot should exist after two turns" + roles = [m.get("role") for m in snapshot.messages] + assert roles == ["user", "assistant", "user", "assistant"], ( + f"Snapshot must contain full transcript for UI hydration, got: {roles}" + ) From 7d51557ce79cad088bea8221d90f0ffa4c76e7ac Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Thu, 20 Aug 2026 00:34:49 +0530 Subject: [PATCH 2/3] fix(python): separate provider input from UI snapshot in AG-UI service-session mode --- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 11 ++++++----- python/packages/ag-ui/tests/ag_ui/test_snapshots.py | 6 +++++- 2 files changed, 11 insertions(+), 6 deletions(-) diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py index 1e84ef6d7bd..8799b9cb607 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py @@ -2272,13 +2272,18 @@ async def run_agent_stream( if not config.use_service_session: raw_messages = snapshot_session.resume_seeded_messages(raw_messages) - + else: + stored_count = len(stored_snapshot.messages) + raw_messages = raw_messages[stored_count:] elif not config.use_service_session: raw_messages = _reconstruct_messages_from_thread_snapshot( stored_messages=stored_snapshot.messages, incoming_messages=raw_messages, stored_interrupt=stored_snapshot.interrupt, ) + else: + stored_count = len(stored_snapshot.messages) + raw_messages = raw_messages[stored_count:] # Initialize flow state with stored state plus request-provided overrides; # endpoint-deferred defaults apply only to keys missing from both. @@ -2853,10 +2858,6 @@ async def run_agent_stream( ) latest_messages_snapshot = snapshot_messages - if config.use_service_session and stored_snapshot is not None and not should_emit_snapshot: - latest_messages_snapshot = _event_messages_to_snapshot_dicts( - list(_build_messages_snapshot(flow, snapshot_messages).messages) - ) if should_emit_snapshot: # Always fold this turn's output into the persisted snapshot, even when the # outbound MESSAGES_SNAPSHOT event is suppressed for predictive tools. diff --git a/python/packages/ag-ui/tests/ag_ui/test_snapshots.py b/python/packages/ag-ui/tests/ag_ui/test_snapshots.py index 8ed1b596d21..de3660b8834 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_snapshots.py +++ b/python/packages/ag-ui/tests/ag_ui/test_snapshots.py @@ -267,7 +267,11 @@ async def test_service_session_snapshot_split_authority() -> None: second_turn = { "threadId": "conv_FHA_SESSION", "__ag_ui_snapshot_scope": "split-auth-test", - "messages": [{"id": "u2", "role": "user", "content": "second"}], + "messages": [ + {"id": "u1", "role": "user", "content": "first"}, + {"id": "a1", "role": "assistant", "content": "ACK"}, + {"id": "u2", "role": "user", "content": "second"}, + ], } await _drain(runner, first_turn) From 35527333f272506793183e28f34db129adcedc36 Mon Sep 17 00:00:00 2001 From: pratikwayase Date: Sun, 23 Aug 2026 12:23:23 +0530 Subject: [PATCH 3/3] fix: use identity matching for service-session split and snapshot reconstruction --- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 77 ++++++++++++++- .../ag-ui/tests/ag_ui/test_snapshots.py | 97 ++++++++++++++++++- 2 files changed, 165 insertions(+), 9 deletions(-) diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py index 8799b9cb607..1d3b529f2f2 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py @@ -2195,6 +2195,60 @@ def _safe_serialize_session_continuation_state( return None +def _split_service_session_input( + stored_snapshot_messages: list[dict[str, Any]], + current_turn_messages: list[dict[str, Any]], +) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + """Splits current turn messages into provider suffix and full snapshot messages. + + Uses identity (message ID) matching with a role-based fallback for assistant + messages (whose agent-generated IDs may differ from what the UI relays) to + safely handle stale snapshots, incremental requests, and truncated histories + without relying on brittle length cursors. + + Args: + stored_snapshot_messages: The messages persisted from the previous turn. + current_turn_messages: The incoming messages from the UI for the current turn. + + Returns: + A tuple containing the provider suffix and the reconstructed snapshot messages. + """ + + def _get_msg_id(msg: Mapping[str, Any]) -> str | None: + return msg.get("id") or msg.get("message_id") + + def _get_msg_role(msg: Mapping[str, Any]) -> str | None: + return msg.get("role") + + i = 0 + j = 0 + while i < len(stored_snapshot_messages) and j < len(current_turn_messages): + stored_msg = stored_snapshot_messages[i] + current_msg = current_turn_messages[j] + stored_id = _get_msg_id(stored_msg) + current_id = _get_msg_id(current_msg) + stored_role = _get_msg_role(stored_msg) + current_role = _get_msg_role(current_msg) + + if stored_id and current_id and stored_id == current_id: + i += 1 + j += 1 + elif stored_role == "assistant" and current_role == "assistant": + i += 1 + j += 1 + else: + break + + provider_suffix = current_turn_messages[j:] + + if j > 0: + snapshot_messages = list(current_turn_messages) + else: + snapshot_messages = list(stored_snapshot_messages) + list(current_turn_messages) + + return provider_suffix, snapshot_messages + + async def run_agent_stream( input_data: dict[str, Any], agent: SupportsAgentRun, @@ -2266,6 +2320,8 @@ async def run_agent_stream( yield event return + snapshot_seed_messages: list[dict[str, Any]] | None = None + if stored_snapshot is not None: if resume_payload is not None and stored_pending_approval_interrupt_ids: seeded_resume_from_snapshot = True @@ -2273,8 +2329,11 @@ async def run_agent_stream( if not config.use_service_session: raw_messages = snapshot_session.resume_seeded_messages(raw_messages) else: - stored_count = len(stored_snapshot.messages) - raw_messages = raw_messages[stored_count:] + provider_suffix, snapshot_seed_messages = _split_service_session_input( + stored_snapshot_messages=stored_snapshot.messages, + current_turn_messages=raw_messages, + ) + raw_messages = provider_suffix elif not config.use_service_session: raw_messages = _reconstruct_messages_from_thread_snapshot( stored_messages=stored_snapshot.messages, @@ -2282,8 +2341,11 @@ async def run_agent_stream( stored_interrupt=stored_snapshot.interrupt, ) else: - stored_count = len(stored_snapshot.messages) - raw_messages = raw_messages[stored_count:] + provider_suffix, snapshot_seed_messages = _split_service_session_input( + stored_snapshot_messages=stored_snapshot.messages, + current_turn_messages=raw_messages, + ) + raw_messages = provider_suffix # Initialize flow state with stored state plus request-provided overrides; # endpoint-deferred defaults apply only to keys missing from both. @@ -2386,7 +2448,12 @@ async def run_agent_stream( protected_tool_call_ids=protected_tool_call_ids, ) - if config.use_service_session and stored_snapshot is not None: + if config.use_service_session and snapshot_seed_messages is not None: + _, snapshot_messages = normalize_agui_input_messages( + snapshot_seed_messages, + protected_tool_call_ids=protected_tool_call_ids, + ) + elif config.use_service_session and stored_snapshot is not None: if seeded_resume_from_snapshot: snapshot_messages = snapshot_session.resume_seeded_messages(snapshot_messages) else: diff --git a/python/packages/ag-ui/tests/ag_ui/test_snapshots.py b/python/packages/ag-ui/tests/ag_ui/test_snapshots.py index de3660b8834..0abf01c230e 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_snapshots.py +++ b/python/packages/ag-ui/tests/ag_ui/test_snapshots.py @@ -2,10 +2,12 @@ """Tests for AG-UI thread snapshot storage primitives.""" +from collections.abc import AsyncGenerator from dataclasses import fields +from typing import Any, cast import pytest -from agent_framework import AgentResponseUpdate, Content +from agent_framework import AgentResponseUpdate, Content, SupportsAgentRun from agent_framework_ag_ui import ( AgentFrameworkAgent, @@ -215,6 +217,8 @@ async def test_in_memory_snapshot_store_rejects_invalid_keys() -> None: class _InputSpyAgent: + """Minimal agent implementation for spying on provider input.""" + name = "spy" description = "" default_options: dict = {} @@ -223,8 +227,8 @@ class _InputSpyAgent: def __init__(self) -> None: self.calls: list[dict] = [] - def run(self, messages, *, session, stream=False, **kwargs): - async def updates(): + def run(self, messages: Any, *, session: Any, stream: bool = False, **kwargs: Any) -> Any: + async def updates() -> AsyncGenerator[Any, None]: self.calls.append( { "roles": [m.role for m in messages], @@ -254,7 +258,7 @@ async def test_service_session_snapshot_split_authority() -> None: agent = _InputSpyAgent() store = InMemoryAGUIThreadSnapshotStore() runner = AgentFrameworkAgent( - agent=agent, + agent=cast(SupportsAgentRun, agent), use_service_session=True, snapshot_store=store, ) @@ -288,3 +292,88 @@ async def test_service_session_snapshot_split_authority() -> None: assert roles == ["user", "assistant", "user", "assistant"], ( f"Snapshot must contain full transcript for UI hydration, got: {roles}" ) + + +@pytest.mark.asyncio +async def test_service_session_snapshot_incremental_request() -> None: + """Verify identity matching handles incremental requests where UI sends only new messages. + + - Stored snapshot has [u1, a1] + - UI sends ONLY [u2] (incremental, not full history) + - Provider should receive [u2] + - Snapshot store should retain [u1, a1, u2, a2] + """ + agent = _InputSpyAgent() + store = InMemoryAGUIThreadSnapshotStore() + runner = AgentFrameworkAgent( + agent=cast(SupportsAgentRun, agent), + use_service_session=True, + snapshot_store=store, + ) + + first_turn = { + "threadId": "conv_INCREMENTAL", + "__ag_ui_snapshot_scope": "incremental-test", + "messages": [{"id": "u1", "role": "user", "content": "first"}], + } + await _drain(runner, first_turn) + + incremental_turn = { + "threadId": "conv_INCREMENTAL", + "__ag_ui_snapshot_scope": "incremental-test", + "messages": [{"id": "u2", "role": "user", "content": "second"}], + } + await _drain(runner, incremental_turn) + + assert agent.calls[1]["roles"] == ["user"], ( + f"Provider should receive incremental input, got: {agent.calls[1]['roles']}" + ) + + snapshot = await store.get(scope="incremental-test", thread_id="conv_INCREMENTAL") + assert snapshot is not None + roles = [m.get("role") for m in snapshot.messages] + assert roles == ["user", "assistant", "user", "assistant"], f"Snapshot must contain full transcript, got: {roles}" + + +@pytest.mark.asyncio +async def test_service_session_snapshot_stale_snapshot() -> None: + """Verify identity matching handles stale snapshots where UI and store diverge. + + - Stored snapshot has [u1, a1] + - UI sends [u1, u2] (UI is ahead, missing a1) + - Provider should receive [u2] (identity match skips u1) + - Snapshot store should retain UI's view [u1, u2, a2] + """ + agent = _InputSpyAgent() + store = InMemoryAGUIThreadSnapshotStore() + runner = AgentFrameworkAgent( + agent=cast(SupportsAgentRun, agent), + use_service_session=True, + snapshot_store=store, + ) + + first_turn = { + "threadId": "conv_STALE", + "__ag_ui_snapshot_scope": "stale-test", + "messages": [{"id": "u1", "role": "user", "content": "first"}], + } + await _drain(runner, first_turn) + + stale_turn = { + "threadId": "conv_STALE", + "__ag_ui_snapshot_scope": "stale-test", + "messages": [ + {"id": "u1", "role": "user", "content": "first"}, + {"id": "u2", "role": "user", "content": "second"}, + ], + } + await _drain(runner, stale_turn) + + assert agent.calls[1]["roles"] == ["user"], ( + f"Provider should receive only new messages, got: {agent.calls[1]['roles']}" + ) + + snapshot = await store.get(scope="stale-test", thread_id="conv_STALE") + assert snapshot is not None + msg_ids = [m.get("id") for m in snapshot.messages] + assert "u1" in msg_ids and "u2" in msg_ids, f"Snapshot should contain UI's messages, got: {msg_ids}"