diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py b/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py index 33d3bede30..6c8e106614 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py @@ -4,9 +4,11 @@ from __future__ import annotations +import json import logging import uuid -from collections.abc import AsyncGenerator, Callable +from collections import Counter +from collections.abc import AsyncGenerator, Callable, Mapping from typing import Any, cast from ag_ui.core import ( @@ -67,6 +69,151 @@ _CHECKPOINT_REQUEST_OWNER_KEY = "ag_ui_workflow_request_owner" +def _hashable_message_content(content: Any) -> Any: + """Return a hashable, order-stable form of snapshot message content.""" + if isinstance(content, (str, int, float, bool)) or content is None: + return content + try: + return json.dumps(content, sort_keys=True, default=str) + except TypeError: + return repr(content) + + +def _pending_request_is_approval(pending_request: Any | None) -> bool: + """Whether a pending request_info event is an approval gate (not conversational HITL).""" + if pending_request is None: + return False + response_type: Any | None + try: + response_type = pending_request.response_type + except Exception: + response_type = getattr(pending_request, "_response_type", None) + if response_type is bool: + return True + type_name = getattr(response_type, "__name__", "") or str(response_type or "") + if "Approval" in type_name: + return True + data = getattr(pending_request, "data", None) + if isinstance(data, dict) and any(key in data for key in ("functionCall", "function_call")): + return True + return False + + +def _snapshot_messages_from_resume_value( + value: Any, + *, + pending_request: Any | None = None, +) -> list[dict[str, Any]]: + """Convert a resolved workflow resume value into snapshot chat messages when user-visible. + + Approval / structured tool payloads are skipped based on the matched pending request's + type/data (not the response text). Only ``user`` turns are projected so resume cannot + forge assistant/system/tool history into the backend-owned snapshot. + """ + if isinstance(value, bool) or value is None: + return [] + if _pending_request_is_approval(pending_request): + return [] + if isinstance(value, str): + text = value.strip() + if not text: + return [] + return [{"role": "user", "content": text}] + if isinstance(value, dict): + # Function-approval style payloads are not chat turns. + if any(key in value for key in ("approved", "accepted", "functionCall", "function_call")): + return [] + if value.get("role") == "user": + return agui_messages_to_snapshot_format([_resume_message_to_agui_dict(value)]) + return [] + if isinstance(value, list): + message_like = [ + _resume_message_to_agui_dict(item) + for item in value + if isinstance(item, dict) and item.get("role") == "user" + ] + if message_like: + return agui_messages_to_snapshot_format(message_like) + return [] + + +def _resume_message_to_agui_dict(message: dict[str, Any]) -> dict[str, Any]: + """Normalize resume message shapes (``contents`` or ``content``) for snapshot encoding.""" + normalized = dict(message) + if normalized.get("content") not in (None, ""): + return normalized + contents = normalized.get("contents") + if isinstance(contents, list): + texts: list[str] = [] + for part in contents: + if isinstance(part, dict) and part.get("type") in {"text", "input_text"}: + texts.append(str(part.get("text") or "")) + if texts: + normalized["content"] = "".join(texts) + return normalized + + +def _snapshot_messages_from_workflow_resume( + resume_payload: Any, + pending_events: Mapping[str, Any] | None = None, +) -> list[dict[str, Any]]: + """Collect user-visible snapshot messages from a workflow resume payload.""" + messages: list[dict[str, Any]] = [] + pending = pending_events or {} + for interrupt in _normalize_resume_interrupts(resume_payload): + if interrupt.get("status") not in {None, "resolved"}: + continue + interrupt_id = interrupt.get("id") + pending_request = pending.get(str(interrupt_id)) if interrupt_id is not None else None + messages.extend( + _snapshot_messages_from_resume_value(interrupt.get("value"), pending_request=pending_request) + ) + return messages + + +def _message_identity(message: dict[str, Any]) -> tuple[Any, ...]: + """Stable identity for deduping resume-synthesized turns against request messages.""" + return (message.get("role"), message.get("id"), _hashable_message_content(message.get("content"))) + + +def _message_content_identity(message: dict[str, Any]) -> tuple[Any, ...]: + """Role+content identity used when message IDs differ across messages vs resume.""" + return (message.get("role"), _hashable_message_content(message.get("content"))) + + +def _append_unique_snapshot_messages( + existing: list[dict[str, Any]], + incoming: list[dict[str, Any]], + *, + content_dedupe_against: list[dict[str, Any]] | None = None, +) -> list[dict[str, Any]]: + """Append resume-derived turns that are not already present in the seed. + + Prefer id equality against ``existing``. Role/content fallback is limited to + ``content_dedupe_against`` (current-turn client overlap / id remaps). Callers + must not pass replayed prior transcript rows here — those keep their ids and + would otherwise let an earlier user ``"yes"`` consume a later resume interrupt + with the same text. When the list is empty or omitted, identical replies across + separate HITL turns on ``messages: []`` resumes are kept. + """ + seen_ids = {message.get("id") for message in existing if message.get("id")} + content_source = content_dedupe_against if content_dedupe_against is not None else [] + remaining_content = Counter(_message_content_identity(message) for message in content_source) + merged = list(existing) + for message in incoming: + message_id = message.get("id") + if message_id and message_id in seen_ids: + continue + content_key = _message_content_identity(message) + if remaining_content[content_key] > 0: + remaining_content[content_key] -= 1 + continue + if message_id: + seen_ids.add(message_id) + merged.append(message) + return merged + + def _checkpoint_id_from_input(input_data: dict[str, Any]) -> str | None: """Read an optional checkpoint id to resume from out of the AG-UI forwarded props.""" forwarded_props = input_data.get("forwarded_props") or input_data.get("forwardedProps") @@ -543,6 +690,9 @@ async def run(self, input_data: dict[str, Any]) -> AsyncGenerator[BaseEvent]: run_id = str(input_data.get("run_id") or input_data.get("runId") or uuid.uuid4()) snapshot_scope = cast(str | None, input_data.get(_SNAPSHOT_SCOPE_INPUT_KEY)) raw_messages = list(cast(list[dict[str, Any]], input_data.get("messages", []) or [])) + # Preserve the client-supplied transcript for content-only dedupe of HITL + # resume turns. Stored history is not a confirmed client-replay overlap. + client_request_messages = list(raw_messages) resume_payload = _extract_resume_payload(input_data) snapshot_session = await ThreadSnapshotSession.open( store=self.snapshot_store, @@ -656,6 +806,34 @@ async def run(self, input_data: dict[str, Any]) -> AsyncGenerator[BaseEvent]: ) else: builder_seed_messages = snapshot_session.resume_seeded_messages(builder_seed_messages) + if resume_payload is not None and snapshot_session.enabled: + # Conversational HITL resumes put the user reply in interrupt.value with + # messages:[]; fold that text into the snapshot so hydrate keeps it (#8160). + # Skip when the client already included the same turn in `messages`. + hitl_messages = _snapshot_messages_from_workflow_resume( + resume_payload, + pending_events=live_pending_events, + ) + if hitl_messages: + # Content fallback is only for the current request's newly supplied + # turns (e.g. client id remap of the resume reply). Rows already in + # the stored snapshot keep their ids when replayed in ``messages`` and + # must not starve a later identical resume interrupt. + stored_ids = { + message.get("id") + for message in (stored_snapshot.messages if stored_snapshot is not None else []) + if message.get("id") + } + current_turn_client_messages = [ + message + for message in client_request_messages + if not (message.get("id") and message.get("id") in stored_ids) + ] + builder_seed_messages = _append_unique_snapshot_messages( + builder_seed_messages, + hitl_messages, + content_dedupe_against=current_turn_client_messages, + ) snapshot_builder = _WorkflowSnapshotBuilder(builder_seed_messages) if snapshot_session.enabled else None if snapshot_builder is not None and effective_state: # Seed builder state so a run that emits no StateSnapshotEvent still diff --git a/python/packages/ag-ui/tests/ag_ui/test_workflow_agent.py b/python/packages/ag-ui/tests/ag_ui/test_workflow_agent.py index 615d3ce4f1..6dfd303940 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_workflow_agent.py +++ b/python/packages/ag-ui/tests/ag_ui/test_workflow_agent.py @@ -10,16 +10,38 @@ from agent_framework import ( Executor, InMemoryCheckpointStorage, + Message, Workflow, WorkflowBuilder, WorkflowContext, executor, handler, + response_handler, ) from agent_framework_ag_ui import AgentFrameworkWorkflow +class _HitlMessageRequestExecutor(Executor): + """Minimal HITL executor shared by snapshot resume regression tests.""" + + def __init__(self) -> None: + super().__init__(id="message_request_executor") + + @handler + async def start(self, message: Any, ctx: WorkflowContext[Any, str]) -> None: + del message + await ctx.request_info({"prompt": "Need user follow-up"}, list[Message], request_id="handoff-user-input") + + @response_handler + async def handle_user_input( + self, original_request: dict[str, Any], response: list[Message], ctx: WorkflowContext[Any, str] + ) -> None: + del original_request + user_text = response[0].text if response else "" + await ctx.yield_output(f"Captured response: {user_text}") + + async def _run(agent: AgentFrameworkWorkflow, payload: dict[str, Any]) -> list[Any]: return [event async for event in agent.run(payload)] @@ -453,3 +475,345 @@ async def finalizer(message: str, ctx: WorkflowContext[None, str]) -> None: # The final assistant text is still preserved alongside it. assert any(message.get("content") == "Final answer." for message in snapshot.messages) + +async def test_workflow_hitl_resume_persists_user_text_in_thread_snapshot() -> None: + """HITL resume with messages:[] must still record the user reply in the snapshot (#8160).""" + from agent_framework_ag_ui import InMemoryAGUIThreadSnapshotStore + from agent_framework_ag_ui._snapshots import _SNAPSHOT_SCOPE_INPUT_KEY, AGUIThreadSnapshot + + storage = InMemoryCheckpointStorage() + workflow = WorkflowBuilder(start_executor=_HitlMessageRequestExecutor()).build() + store = InMemoryAGUIThreadSnapshotStore() + agent = AgentFrameworkWorkflow(workflow=workflow, snapshot_store=store, checkpoint_storage=storage) + + first_events = await _run( + agent, + { + "thread_id": "thread-hitl", + "run_id": "run-1", + "messages": [{"id": "user-1", "role": "user", "content": "start"}], + _SNAPSHOT_SCOPE_INPUT_KEY: "tenant-a", + }, + ) + assert "RUN_ERROR" not in [event.type for event in first_events] + + await store.save( + scope="tenant-a", + thread_id="thread-hitl", + snapshot=AGUIThreadSnapshot( + messages=[ + {"id": "user-1", "role": "user", "content": "start"}, + {"id": "assistant-1", "role": "assistant", "content": "Need more detail"}, + ], + state=None, + interrupt=None, + ), + ) + + resumed_events = await _run( + agent, + { + "thread_id": "thread-hitl", + "run_id": "run-2", + "messages": [], + "resume": { + "interrupts": [ + { + "id": "handoff-user-input", + "value": [ + { + "role": "user", + "contents": [{"type": "text", "text": "Please ship a replacement instead."}], + } + ], + } + ] + }, + _SNAPSHOT_SCOPE_INPUT_KEY: "tenant-a", + }, + ) + assert "RUN_ERROR" not in [event.type for event in resumed_events] + + snapshot = await store.get(scope="tenant-a", thread_id="thread-hitl") + assert snapshot is not None + contents = [message.get("content") for message in snapshot.messages] + assert "start" in contents + assert any(isinstance(content, str) and "replacement" in content for content in contents) + assert any( + message.get("role") == "user" + and isinstance(message.get("content"), str) + and "replacement" in message["content"] + for message in snapshot.messages + ) + + +async def test_workflow_hitl_resume_keeps_repeated_yes_on_empty_messages() -> None: + """A second HITL 'yes' with messages:[] must not be dropped as a content duplicate.""" + from agent_framework_ag_ui import InMemoryAGUIThreadSnapshotStore + from agent_framework_ag_ui._snapshots import _SNAPSHOT_SCOPE_INPUT_KEY, AGUIThreadSnapshot + + storage = InMemoryCheckpointStorage() + workflow = WorkflowBuilder(start_executor=_HitlMessageRequestExecutor()).build() + store = InMemoryAGUIThreadSnapshotStore() + agent = AgentFrameworkWorkflow(workflow=workflow, snapshot_store=store, checkpoint_storage=storage) + + # Establish a live pending request_info interrupt on this workflow instance. + first_events = await _run( + agent, + { + "thread_id": "thread-hitl-yes", + "run_id": "run-yes-1", + "messages": [{"id": "user-1", "role": "user", "content": "start"}], + _SNAPSHOT_SCOPE_INPUT_KEY: "tenant-a", + }, + ) + assert "RUN_ERROR" not in [event.type for event in first_events] + + # Simulate a thread that already persisted an earlier identical "yes". + await store.save( + scope="tenant-a", + thread_id="thread-hitl-yes", + snapshot=AGUIThreadSnapshot( + messages=[ + {"id": "user-1", "role": "user", "content": "start"}, + {"id": "assistant-1", "role": "assistant", "content": "confirm?"}, + {"id": "user-yes-1", "role": "user", "content": "yes"}, + {"id": "assistant-2", "role": "assistant", "content": "confirm again?"}, + ], + state=None, + interrupt=None, + ), + ) + + resumed_events = await _run( + agent, + { + "thread_id": "thread-hitl-yes", + "run_id": "run-yes-2", + "messages": [], + "resume": { + "interrupts": [ + { + "id": "handoff-user-input", + "value": [ + { + "role": "user", + "contents": [{"type": "text", "text": "yes"}], + } + ], + } + ] + }, + _SNAPSHOT_SCOPE_INPUT_KEY: "tenant-a", + }, + ) + assert "RUN_ERROR" not in [event.type for event in resumed_events] + + snapshot = await store.get(scope="tenant-a", thread_id="thread-hitl-yes") + assert snapshot is not None + yes_turns = [ + message + for message in snapshot.messages + if message.get("role") == "user" and message.get("content") == "yes" + ] + assert len(yes_turns) >= 2 + + +async def test_workflow_hitl_resume_keeps_yes_when_messages_replay_prior_yes() -> None: + """Client-replayed prior 'yes' in messages must not drop a new resume interrupt yes.""" + from agent_framework_ag_ui import InMemoryAGUIThreadSnapshotStore + from agent_framework_ag_ui._snapshots import _SNAPSHOT_SCOPE_INPUT_KEY, AGUIThreadSnapshot + + storage = InMemoryCheckpointStorage() + workflow = WorkflowBuilder(start_executor=_HitlMessageRequestExecutor()).build() + store = InMemoryAGUIThreadSnapshotStore() + agent = AgentFrameworkWorkflow(workflow=workflow, snapshot_store=store, checkpoint_storage=storage) + + prior = [ + {"id": "user-1", "role": "user", "content": "start"}, + {"id": "assistant-1", "role": "assistant", "content": "confirm?"}, + {"id": "user-yes-1", "role": "user", "content": "yes"}, + {"id": "assistant-2", "role": "assistant", "content": "confirm again?"}, + ] + + first_events = await _run( + agent, + { + "thread_id": "thread-hitl-replay-yes", + "run_id": "run-yes-setup", + "messages": [{"id": "user-1", "role": "user", "content": "start"}], + _SNAPSHOT_SCOPE_INPUT_KEY: "tenant-a", + }, + ) + assert "RUN_ERROR" not in [event.type for event in first_events] + + await store.save( + scope="tenant-a", + thread_id="thread-hitl-replay-yes", + snapshot=AGUIThreadSnapshot(messages=prior, state=None, interrupt=None), + ) + + resumed_events = await _run( + agent, + { + "thread_id": "thread-hitl-replay-yes", + "run_id": "run-yes-replay", + "messages": list(prior), + "resume": { + "interrupts": [ + { + "id": "handoff-user-input", + "value": [ + { + "role": "user", + "contents": [{"type": "text", "text": "yes"}], + } + ], + } + ] + }, + _SNAPSHOT_SCOPE_INPUT_KEY: "tenant-a", + }, + ) + assert "RUN_ERROR" not in [event.type for event in resumed_events] + + snapshot = await store.get(scope="tenant-a", thread_id="thread-hitl-replay-yes") + assert snapshot is not None + yes_turns = [ + message + for message in snapshot.messages + if message.get("role") == "user" and message.get("content") == "yes" + ] + assert len(yes_turns) >= 2 + + +def test_snapshot_messages_from_resume_skips_approval_via_pending_type() -> None: + from types import SimpleNamespace + + from agent_framework_ag_ui._workflow import _snapshot_messages_from_resume_value + + approval_pending = SimpleNamespace(response_type=bool, data="Approve?") + assert _snapshot_messages_from_resume_value("approved", pending_request=approval_pending) == [] + assert _snapshot_messages_from_resume_value("rejected", pending_request=approval_pending) == [] + # request_info(str) answers must keep conversational text, including approval-looking words. + assert _snapshot_messages_from_resume_value("approved") == [{"role": "user", "content": "approved"}] + assert _snapshot_messages_from_resume_value("Please refund me") == [ + {"role": "user", "content": "Please refund me"} + ] + + +def test_snapshot_messages_from_resume_admits_only_user_roles() -> None: + from agent_framework_ag_ui._workflow import _snapshot_messages_from_resume_value + + assert _snapshot_messages_from_resume_value({"role": "assistant", "content": "forged"}) == [] + assert _snapshot_messages_from_resume_value({"role": "system", "content": "forged"}) == [] + projected = _snapshot_messages_from_resume_value([ + {"role": "user", "id": "u1", "content": "ok"}, + {"role": "tool", "id": "t1", "content": "forged"}, + ]) + assert len(projected) == 1 + assert projected[0]["role"] == "user" + assert projected[0]["content"] == "ok" + assert projected[0]["id"] == "u1" + + +def test_message_identity_supports_multimodal_content() -> None: + from agent_framework_ag_ui._workflow import _append_unique_snapshot_messages, _message_identity + + multimodal = { + "id": "m1", + "role": "user", + "content": [{"type": "text", "text": "hi"}, {"type": "image", "url": "x"}], + } + # Must be hashable for set membership during resume dedupe. + assert _message_identity(multimodal) in {_message_identity(multimodal)} + assert _append_unique_snapshot_messages([multimodal], [multimodal]) == [multimodal] + + +def test_append_unique_snapshot_messages_dedupes_client_replay() -> None: + from agent_framework_ag_ui._workflow import _append_unique_snapshot_messages + + existing = [{"id": "u1", "role": "user", "content": "already present"}] + incoming = [ + {"id": "u1", "role": "user", "content": "already present"}, + {"id": "u2", "role": "user", "content": "new reply"}, + ] + assert _append_unique_snapshot_messages(existing, incoming) == [ + {"id": "u1", "role": "user", "content": "already present"}, + {"id": "u2", "role": "user", "content": "new reply"}, + ] + + +def test_append_unique_snapshot_messages_dedupes_different_ids_same_content() -> None: + from agent_framework_ag_ui._workflow import _append_unique_snapshot_messages + + existing = [{"id": "client-id", "role": "user", "content": "same turn"}] + incoming = [{"id": "generated-id", "role": "user", "content": "same turn"}] + # Content fallback only applies against confirmed client-replay overlap. + assert ( + _append_unique_snapshot_messages( + existing, + incoming, + content_dedupe_against=existing, + ) + == existing + ) + + +def test_append_unique_snapshot_messages_keeps_intentional_repeated_replies() -> None: + from agent_framework_ag_ui._workflow import _append_unique_snapshot_messages + + existing = [{"id": "u0", "role": "user", "content": "hello"}] + incoming = [ + {"id": "r1", "role": "user", "content": "repeat"}, + {"id": "r2", "role": "user", "content": "repeat"}, + ] + merged = _append_unique_snapshot_messages(existing, incoming) + assert [m["id"] for m in merged] == ["u0", "r1", "r2"] + + +def test_append_unique_snapshot_messages_keeps_second_hitl_yes_without_client_replay() -> None: + """messages:[] HITL resumes must not collapse a later identical reply against history.""" + from agent_framework_ag_ui._workflow import _append_unique_snapshot_messages + + history = [ + {"id": "u0", "role": "user", "content": "start"}, + {"id": "a0", "role": "assistant", "content": "confirm?"}, + {"id": "u1", "role": "user", "content": "yes"}, + {"id": "a1", "role": "assistant", "content": "confirm again?"}, + ] + second_yes = [{"id": "generated-yes-2", "role": "user", "content": "yes"}] + merged = _append_unique_snapshot_messages( + history, + second_yes, + content_dedupe_against=[], + ) + assert [m["id"] for m in merged] == ["u0", "a0", "u1", "a1", "generated-yes-2"] + + +def test_append_unique_snapshot_messages_keeps_resume_yes_when_client_replays_prior_yes() -> None: + """Replayed prior-yes rows (same ids as stored) must not be passed as content fallback.""" + from agent_framework_ag_ui._workflow import _append_unique_snapshot_messages + + history = [ + {"id": "u0", "role": "user", "content": "start"}, + {"id": "a0", "role": "assistant", "content": "confirm?"}, + {"id": "u1", "role": "user", "content": "yes"}, + {"id": "a1", "role": "assistant", "content": "confirm again?"}, + ] + stored_ids = {message.get("id") for message in history if message.get("id")} + client_replay = list(history) + # Mirrors the run() call site: only count client turns not already stored by id. + current_turn_client_messages = [ + message + for message in client_replay + if not (message.get("id") and message.get("id") in stored_ids) + ] + second_yes = [{"id": "generated-yes-2", "role": "user", "content": "yes"}] + merged = _append_unique_snapshot_messages( + history, + second_yes, + content_dedupe_against=current_turn_client_messages, + ) + assert current_turn_client_messages == [] + assert [m["id"] for m in merged] == ["u0", "a0", "u1", "a1", "generated-yes-2"]